Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/search_endpoints/search_tool_registry.py: 76%
105 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Search Tool Registry for managing search tool configurations.
3"""
5from collections.abc import Iterator, Mapping, Sequence
6from datetime import datetime, timezone
7from typing import Final, Protocol
9from litellm._logging import verbose_proxy_logger
10from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
11from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry
12from litellm.proxy.utils import PrismaClient
13from litellm.repositories.table_repositories import SearchToolsRepository
14from litellm.types.search import SearchTool
17class SearchToolRecord(Protocol):
18 search_tool_id: str
19 search_tool_name: str
20 created_at: datetime
21 updated_at: datetime
23 def __iter__(self) -> Iterator[tuple[str, object]]: ... 23 ↛ exitline 23 didn't return from function '__iter__' because
26class SearchToolTableClient(Protocol):
27 async def create(self, data: Mapping[str, object]) -> SearchToolRecord: ... 27 ↛ exitline 27 didn't return from function 'create' because
29 async def find_unique(self, where: Mapping[str, object]) -> SearchToolRecord | None: ... 29 ↛ exitline 29 didn't return from function 'find_unique' because
31 async def find_many(self, order: Mapping[str, str] | None = None) -> Sequence[SearchToolRecord]: ... 31 ↛ exitline 31 didn't return from function 'find_many' because
33 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> SearchToolRecord: ... 33 ↛ exitline 33 didn't return from function 'update' because
35 async def delete(self, where: Mapping[str, object]) -> SearchToolRecord: ... 35 ↛ exitline 35 didn't return from function 'delete' because
38class _SearchToolsRepositoryView(Protocol):
39 @property
40 def table(self) -> SearchToolTableClient: ... 40 ↛ exitline 40 didn't return from function 'table' because
43def _search_tools_table_of(repository: _SearchToolsRepositoryView) -> SearchToolTableClient:
44 return repository.table
47def _search_tools_table(prisma_client: PrismaClient) -> SearchToolTableClient:
48 return _search_tools_table_of(SearchToolsRepository(prisma_client))
51class SearchToolRegistry:
52 """
53 Handles adding, removing, and getting search tools in DB + in memory.
54 """
56 def __init__(self):
57 pass
59 @staticmethod
60 def _convert_prisma_to_dict(prisma_obj: SearchToolRecord) -> dict:
61 """
62 Convert Prisma result to dict with datetime objects as ISO format strings.
64 Args:
65 prisma_obj: Prisma model instance
67 Returns:
68 Dict with datetime fields converted to ISO strings
69 """
70 result: Final = dict(prisma_obj)
71 # Convert datetime objects to ISO format strings
72 if "created_at" in result and result["created_at"]: 72 ↛ 74line 72 didn't jump to line 74 because the condition on line 72 was always true
73 result["created_at"] = prisma_obj.created_at.isoformat()
74 if "updated_at" in result and result["updated_at"]: 74 ↛ 76line 74 didn't jump to line 76 because the condition on line 74 was always true
75 result["updated_at"] = prisma_obj.updated_at.isoformat()
76 return result
78 ###########################################################
79 ########### DB management helpers for search tools ########
80 ###########################################################
82 async def add_search_tool_to_db(self, search_tool: SearchTool, prisma_client: PrismaClient):
83 """
84 Add a search tool to the database.
86 Args:
87 search_tool: Search tool configuration
88 prisma_client: Prisma client instance
90 Returns:
91 Dict with created search tool data
92 """
93 try:
94 search_tool_name: Final = search_tool.get("search_tool_name")
95 litellm_params: Final[str] = safe_dumps(dict(search_tool.get("litellm_params", {})))
96 search_tool_info: Final[str] = safe_dumps(search_tool.get("search_tool_info", {}))
98 # Create search tool in DB
99 created_search_tool: Final = await _search_tools_table(prisma_client).create(
100 data={
101 "search_tool_name": search_tool_name,
102 "litellm_params": litellm_params,
103 "search_tool_info": search_tool_info,
104 "created_at": datetime.now(timezone.utc),
105 "updated_at": datetime.now(timezone.utc),
106 }
107 )
109 # Add search_tool_id to the returned search tool object
110 search_tool_dict: Final = dict(search_tool)
111 search_tool_dict["search_tool_id"] = created_search_tool.search_tool_id
112 search_tool_dict["created_at"] = created_search_tool.created_at.isoformat()
113 search_tool_dict["updated_at"] = created_search_tool.updated_at.isoformat()
115 return search_tool_dict
116 except Exception as e:
117 verbose_proxy_logger.exception("Error adding search tool to DB: %s", e)
118 raise Exception(f"Error adding search tool to DB: {e}")
120 async def delete_search_tool_from_db(self, search_tool_id: str, prisma_client: PrismaClient):
121 """
122 Delete a search tool from the database.
124 Args:
125 search_tool_id: ID of search tool to delete
126 prisma_client: Prisma client instance
128 Returns:
129 Dict with success message
130 """
131 try:
132 # Get search tool before deletion for response
133 existing_tool: Final = await _search_tools_table(prisma_client).find_unique(
134 where={"search_tool_id": search_tool_id}
135 )
137 if not existing_tool: 137 ↛ 138line 137 didn't jump to line 138 because the condition on line 137 was never true
138 raise Exception(f"Search tool with ID {search_tool_id} not found")
140 # Delete from DB
141 await _search_tools_table(prisma_client).delete(where={"search_tool_id": search_tool_id})
143 return {
144 "message": f"Search tool {search_tool_id} deleted successfully",
145 "search_tool_name": existing_tool.search_tool_name,
146 }
147 except Exception as e:
148 verbose_proxy_logger.exception("Error deleting search tool from DB: %s", e)
149 raise Exception(f"Error deleting search tool from DB: {e}")
151 async def update_search_tool_in_db(self, search_tool_id: str, search_tool: SearchTool, prisma_client: PrismaClient):
152 """
153 Update a search tool in the database.
155 Args:
156 search_tool_id: ID of search tool to update
157 search_tool: Updated search tool configuration
158 prisma_client: Prisma client instance
160 Returns:
161 Dict with updated search tool data
162 """
163 try:
164 search_tool_name: Final = search_tool.get("search_tool_name")
165 litellm_params: Final[str] = safe_dumps(dict(search_tool.get("litellm_params", {})))
166 search_tool_info: Final[str] = safe_dumps(search_tool.get("search_tool_info", {}))
168 # Update in DB
169 updated_search_tool: Final = await _search_tools_table(prisma_client).update(
170 where={"search_tool_id": search_tool_id},
171 data={
172 "search_tool_name": search_tool_name,
173 "litellm_params": litellm_params,
174 "search_tool_info": search_tool_info,
175 "updated_at": datetime.now(timezone.utc),
176 },
177 )
179 # Convert to dict with ISO formatted datetimes
180 return self._convert_prisma_to_dict(updated_search_tool)
181 except Exception as e:
182 verbose_proxy_logger.exception("Error updating search tool in DB: %s", e)
183 raise Exception(f"Error updating search tool in DB: {e}")
185 @staticmethod
186 async def get_all_search_tools_from_db(
187 prisma_client: PrismaClient,
188 ) -> list[SearchTool]:
189 """
190 Get all search tools from the database.
192 Args:
193 prisma_client: Prisma client instance
195 Returns:
196 List of search tool configurations
197 """
198 try:
199 search_tools_from_db: Final = await call_with_db_reconnect_retry(
200 prisma_client,
201 lambda: _search_tools_table(prisma_client).find_many(
202 order={"created_at": "desc"},
203 ),
204 reason="get_all_search_tools_from_db_lookup_failure",
205 )
207 search_tools: Final[list[SearchTool]] = []
208 for search_tool in search_tools_from_db:
209 # Convert Prisma result to dict with ISO formatted datetimes
210 search_tool_dict = SearchToolRegistry._convert_prisma_to_dict(search_tool)
211 search_tools.append(SearchTool(**search_tool_dict))
213 return search_tools
214 except Exception as e:
215 verbose_proxy_logger.exception("Error getting search tools from DB: %s", e)
216 raise Exception(f"Error getting search tools from DB: {e}")
218 async def get_search_tool_by_id_from_db(
219 self, search_tool_id: str, prisma_client: PrismaClient
220 ) -> SearchTool | None:
221 """
222 Get a search tool by its ID from the database.
224 Args:
225 search_tool_id: ID of search tool to retrieve
226 prisma_client: Prisma client instance
228 Returns:
229 Search tool configuration or None if not found
230 """
231 try:
232 search_tool: Final = await _search_tools_table(prisma_client).find_unique(
233 where={"search_tool_id": search_tool_id}
234 )
236 if not search_tool:
237 return None
239 # Convert Prisma result to dict with ISO formatted datetimes
240 search_tool_dict: Final = self._convert_prisma_to_dict(search_tool)
241 return SearchTool(**search_tool_dict)
242 except Exception as e:
243 verbose_proxy_logger.exception("Error getting search tool from DB: %s", e)
244 raise Exception(f"Error getting search tool from DB: {e}")
246 async def get_search_tool_by_name_from_db(
247 self, search_tool_name: str, prisma_client: PrismaClient
248 ) -> SearchTool | None:
249 """
250 Get a search tool by its name from the database.
252 Args:
253 search_tool_name: Name of search tool to retrieve
254 prisma_client: Prisma client instance
256 Returns:
257 Search tool configuration or None if not found
258 """
259 try:
260 search_tool: Final = await _search_tools_table(prisma_client).find_unique(
261 where={"search_tool_name": search_tool_name}
262 )
264 if not search_tool:
265 return None
267 # Convert Prisma result to dict with ISO formatted datetimes
268 search_tool_dict: Final = self._convert_prisma_to_dict(search_tool)
269 return SearchTool(**search_tool_dict)
270 except Exception as e:
271 verbose_proxy_logger.exception("Error getting search tool from DB: %s", e)
272 raise Exception(f"Error getting search tool from DB: {e}")