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

1""" 

2Search Tool Registry for managing search tool configurations. 

3""" 

4 

5from collections.abc import Iterator, Mapping, Sequence 

6from datetime import datetime, timezone 

7from typing import Final, Protocol 

8 

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 

15 

16 

17class SearchToolRecord(Protocol): 

18 search_tool_id: str 

19 search_tool_name: str 

20 created_at: datetime 

21 updated_at: datetime 

22 

23 def __iter__(self) -> Iterator[tuple[str, object]]: ... 23 ↛ exitline 23 didn't return from function '__iter__' because

24 

25 

26class SearchToolTableClient(Protocol): 

27 async def create(self, data: Mapping[str, object]) -> SearchToolRecord: ... 27 ↛ exitline 27 didn't return from function 'create' because

28 

29 async def find_unique(self, where: Mapping[str, object]) -> SearchToolRecord | None: ... 29 ↛ exitline 29 didn't return from function 'find_unique' because

30 

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

32 

33 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> SearchToolRecord: ... 33 ↛ exitline 33 didn't return from function 'update' because

34 

35 async def delete(self, where: Mapping[str, object]) -> SearchToolRecord: ... 35 ↛ exitline 35 didn't return from function 'delete' because

36 

37 

38class _SearchToolsRepositoryView(Protocol): 

39 @property 

40 def table(self) -> SearchToolTableClient: ... 40 ↛ exitline 40 didn't return from function 'table' because

41 

42 

43def _search_tools_table_of(repository: _SearchToolsRepositoryView) -> SearchToolTableClient: 

44 return repository.table 

45 

46 

47def _search_tools_table(prisma_client: PrismaClient) -> SearchToolTableClient: 

48 return _search_tools_table_of(SearchToolsRepository(prisma_client)) 

49 

50 

51class SearchToolRegistry: 

52 """ 

53 Handles adding, removing, and getting search tools in DB + in memory. 

54 """ 

55 

56 def __init__(self): 

57 pass 

58 

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. 

63 

64 Args: 

65 prisma_obj: Prisma model instance 

66 

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 

77 

78 ########################################################### 

79 ########### DB management helpers for search tools ######## 

80 ########################################################### 

81 

82 async def add_search_tool_to_db(self, search_tool: SearchTool, prisma_client: PrismaClient): 

83 """ 

84 Add a search tool to the database. 

85 

86 Args: 

87 search_tool: Search tool configuration 

88 prisma_client: Prisma client instance 

89 

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", {})) 

97 

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 ) 

108 

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() 

114 

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}") 

119 

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. 

123 

124 Args: 

125 search_tool_id: ID of search tool to delete 

126 prisma_client: Prisma client instance 

127 

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 ) 

136 

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") 

139 

140 # Delete from DB 

141 await _search_tools_table(prisma_client).delete(where={"search_tool_id": search_tool_id}) 

142 

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}") 

150 

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. 

154 

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 

159 

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", {})) 

167 

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 ) 

178 

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}") 

184 

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. 

191 

192 Args: 

193 prisma_client: Prisma client instance 

194 

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 ) 

206 

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)) 

212 

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}") 

217 

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. 

223 

224 Args: 

225 search_tool_id: ID of search tool to retrieve 

226 prisma_client: Prisma client instance 

227 

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 ) 

235 

236 if not search_tool: 

237 return None 

238 

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}") 

245 

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. 

251 

252 Args: 

253 search_tool_name: Name of search tool to retrieve 

254 prisma_client: Prisma client instance 

255 

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 ) 

263 

264 if not search_tool: 

265 return None 

266 

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}")