Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/pass_through_endpoints/llm_provider_handlers/assembly_passthrough_logging_handler.py: 34%

131 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import asyncio 

2import json 

3import time 

4import urllib.parse 

5from datetime import datetime 

6from typing import Final, Literal 

7from urllib.parse import urlparse 

8 

9import httpx 

10from typing_extensions import TypedDict 

11 

12import litellm 

13from litellm._logging import verbose_proxy_logger 

14from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

15from litellm.litellm_core_utils.litellm_logging import ( 

16 get_standard_logging_object_payload, 

17) 

18from litellm.litellm_core_utils.thread_pool_executor import executor 

19from litellm.types.passthrough_endpoints.assembly_ai import ( 

20 ASSEMBLY_AI_MAX_POLLING_ATTEMPTS, 

21 ASSEMBLY_AI_POLLING_INTERVAL, 

22) 

23from litellm.types.passthrough_endpoints.pass_through_endpoints import ( 

24 PassthroughStandardLoggingPayload, 

25) 

26 

27 

28class AssemblyAITranscriptResponse(TypedDict, total=False): 

29 id: str 

30 speech_model: str 

31 acoustic_model: str 

32 language_code: str 

33 status: str 

34 audio_duration: float 

35 

36 

37class AssemblyAIPassthroughLoggingHandler: 

38 def __init__(self): 

39 self.assembly_ai_base_url = "https://api.assemblyai.com" 

40 self.assembly_ai_eu_base_url = "https://eu.assemblyai.com" 

41 """ 

42 The base URL for the AssemblyAI API 

43 """ 

44 

45 self.polling_interval: float = ASSEMBLY_AI_POLLING_INTERVAL 

46 """ 

47 The polling interval for the AssemblyAI API.  

48 litellm needs to poll the GET /transcript/{transcript_id} endpoint to get the status of the transcript. 

49 """ 

50 

51 self.max_polling_attempts = ASSEMBLY_AI_MAX_POLLING_ATTEMPTS 

52 """ 

53 The maximum number of polling attempts for the AssemblyAI API. 

54 """ 

55 

56 def assemblyai_passthrough_logging_handler( 

57 self, 

58 httpx_response: httpx.Response, 

59 response_body: dict, 

60 logging_obj: LiteLLMLoggingObj, 

61 url_route: str, 

62 result: str, 

63 start_time: datetime, 

64 end_time: datetime, 

65 cache_hit: bool, 

66 **kwargs, 

67 ): 

68 """ 

69 Since cost tracking requires polling the AssemblyAI API, we need to handle this in a separate thread. Hence the executor.submit. 

70 """ 

71 executor.submit( 

72 self._handle_assemblyai_passthrough_logging, 

73 httpx_response, 

74 response_body, 

75 logging_obj, 

76 url_route, 

77 result, 

78 start_time, 

79 end_time, 

80 cache_hit, 

81 **kwargs, 

82 ) 

83 

84 def _handle_assemblyai_passthrough_logging( 

85 self, 

86 httpx_response: httpx.Response, 

87 response_body: dict, 

88 logging_obj: LiteLLMLoggingObj, 

89 url_route: str, 

90 result: str, 

91 start_time: datetime, 

92 end_time: datetime, 

93 cache_hit: bool, 

94 **kwargs, 

95 ): 

96 """ 

97 Handles logging for AssemblyAI successful passthrough requests 

98 """ 

99 from ..pass_through_endpoints import pass_through_endpoint_logging 

100 

101 model: Final = response_body.get("speech_model", "") 

102 verbose_proxy_logger.debug("response body %s", json.dumps(response_body, indent=4)) 

103 kwargs["model"] = model 

104 kwargs["custom_llm_provider"] = "assemblyai" 

105 response_cost: float | None = None 

106 

107 transcript_id: Final = response_body.get("id") 

108 if transcript_id is None: 

109 raise ValueError("Transcript ID is required to log the cost of the transcription") 

110 transcript_response: Final = self._poll_assembly_for_transcript_response( 

111 transcript_id=transcript_id, url_route=url_route 

112 ) 

113 verbose_proxy_logger.debug( 

114 "finished polling assembly for transcript response- got transcript response %s", 

115 json.dumps(transcript_response, indent=4), 

116 ) 

117 if transcript_response: 

118 cost: Final = self.get_cost_for_assembly_transcript( 

119 speech_model=model, 

120 transcript_response=transcript_response, 

121 ) 

122 response_cost = cost 

123 

124 # Make standard logging object for Vertex AI 

125 standard_logging_object: Final = get_standard_logging_object_payload( 

126 kwargs=kwargs, 

127 init_response_obj=transcript_response, 

128 start_time=start_time, 

129 end_time=end_time, 

130 logging_obj=logging_obj, 

131 status="success", 

132 ) 

133 

134 passthrough_logging_payload: Final[PassthroughStandardLoggingPayload | None] = kwargs.get( 

135 "passthrough_logging_payload" 

136 ) 

137 

138 verbose_proxy_logger.debug( 

139 "standard_passthrough_logging_object %s", 

140 json.dumps(passthrough_logging_payload, indent=4), 

141 ) 

142 

143 # pretty print standard logging object 

144 verbose_proxy_logger.debug("standard_logging_object= %s", json.dumps(standard_logging_object, indent=4)) 

145 logging_obj.model_call_details["model"] = model 

146 logging_obj.model_call_details["custom_llm_provider"] = "assemblyai" 

147 logging_obj.model_call_details["response_cost"] = response_cost 

148 

149 asyncio.run( 

150 pass_through_endpoint_logging._handle_logging( 

151 logging_obj=logging_obj, 

152 standard_logging_response_object=self._get_response_to_log(transcript_response), 

153 result=result, 

154 start_time=start_time, 

155 end_time=end_time, 

156 cache_hit=cache_hit, 

157 **kwargs, 

158 ) 

159 ) 

160 

161 def _get_response_to_log(self, transcript_response: AssemblyAITranscriptResponse | None) -> dict: 

162 if transcript_response is None: 

163 return {} 

164 return dict(transcript_response) 

165 

166 def _get_assembly_transcript( 

167 self, 

168 transcript_id: str, 

169 request_region: Literal["eu"] | None = None, 

170 ) -> dict | None: 

171 """ 

172 Get the transcript details from AssemblyAI API 

173 

174 Args: 

175 response_body (dict): Response containing the transcript ID 

176 

177 Returns: 

178 Optional[dict]: Transcript details if successful, None otherwise 

179 """ 

180 from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( 

181 passthrough_endpoint_router, 

182 ) 

183 

184 _base_url: Final = self.assembly_ai_eu_base_url if request_region == "eu" else self.assembly_ai_base_url 

185 _api_key: Final = passthrough_endpoint_router.get_credentials( 

186 custom_llm_provider="assemblyai", 

187 region_name=request_region, 

188 ) 

189 if _api_key is None: 

190 raise ValueError("AssemblyAI API key not found") 

191 if any(c in transcript_id for c in ("/", "\\", "#", "?")) or ".." in transcript_id: 

192 raise ValueError(f"Invalid transcript_id {transcript_id!r}: contains disallowed characters") 

193 safe_transcript_id: Final = urllib.parse.quote(transcript_id, safe="") 

194 try: 

195 url: Final = f"{_base_url}/v2/transcript/{safe_transcript_id}" 

196 headers: Final = { 

197 "Authorization": f"Bearer {_api_key}", 

198 "Content-Type": "application/json", 

199 } 

200 

201 response: Final = httpx.get(url, headers=headers) 

202 response.raise_for_status() 

203 

204 return response.json() 

205 except Exception as e: 

206 verbose_proxy_logger.exception("[Non blocking logging error] Error getting AssemblyAI transcript: %s", e) 

207 return None 

208 

209 def _poll_assembly_for_transcript_response( 

210 self, 

211 transcript_id: str, 

212 url_route: str | None = None, 

213 ) -> AssemblyAITranscriptResponse | None: 

214 """ 

215 Poll the status of the transcript until it is completed or timeout (30 minutes) 

216 """ 

217 for _ in range(self.max_polling_attempts): # 180 attempts * 10s = 30 minutes max 

218 transcript = self._get_assembly_transcript( 

219 request_region=AssemblyAIPassthroughLoggingHandler._get_assembly_region_from_url(url=url_route), 

220 transcript_id=transcript_id, 

221 ) 

222 if transcript is None: 

223 return None 

224 if transcript.get("status") == "completed" or transcript.get("status") == "error": 

225 return AssemblyAITranscriptResponse(**transcript) 

226 time.sleep(self.polling_interval) 

227 return None 

228 

229 @staticmethod 

230 def get_cost_for_assembly_transcript( 

231 transcript_response: AssemblyAITranscriptResponse, 

232 speech_model: str, 

233 ) -> float | None: 

234 """ 

235 Get the cost for the assembly transcript 

236 """ 

237 _audio_duration: Final = transcript_response.get("audio_duration") 

238 if _audio_duration is None: 

239 return None 

240 _cost_per_second: Final = AssemblyAIPassthroughLoggingHandler.get_cost_per_second_for_assembly_model( 

241 speech_model=speech_model 

242 ) 

243 if _cost_per_second is None: 

244 return None 

245 return _audio_duration * _cost_per_second 

246 

247 @staticmethod 

248 def get_cost_per_second_for_assembly_model(speech_model: str) -> float | None: 

249 """ 

250 Get the cost per second for the assembly model. 

251 Falls back to assemblyai/nano if the specific speech model info cannot be found. 

252 """ 

253 try: 

254 # First try with the provided speech model 

255 try: 

256 model_info = litellm.get_model_info( 

257 model=speech_model, 

258 custom_llm_provider="assemblyai", 

259 ) 

260 if model_info and model_info.get("input_cost_per_second") is not None: 

261 return model_info.get("input_cost_per_second") 

262 except Exception: 

263 pass # Continue to fallback if model not found 

264 

265 # Fallback to assemblyai/nano if speech model info not found 

266 try: 

267 model_info = litellm.get_model_info( 

268 model="assemblyai/nano", 

269 custom_llm_provider="assemblyai", 

270 ) 

271 if model_info and model_info.get("input_cost_per_second") is not None: 

272 return model_info.get("input_cost_per_second") 

273 except Exception: 

274 pass 

275 

276 return None 

277 except Exception as e: 

278 verbose_proxy_logger.exception("[Non blocking logging error] Error getting AssemblyAI model info: %s", e) 

279 return None 

280 

281 @staticmethod 

282 def _should_log_request(request_method: str) -> bool: 

283 """ 

284 only POST transcription jobs are logged. litellm will POLL assembly to wait for the transcription to complete to log the complete response / cost 

285 """ 

286 return request_method == "POST" 

287 

288 @staticmethod 

289 def _get_assembly_region_from_url(url: str | None) -> Literal["eu"] | None: 

290 """ 

291 Get the region from the URL 

292 """ 

293 if url is None: 293 ↛ 294line 293 didn't jump to line 294 because the condition on line 293 was never true

294 return None 

295 if urlparse(url).hostname == "eu.assemblyai.com": 295 ↛ 296line 295 didn't jump to line 296 because the condition on line 295 was never true

296 return "eu" 

297 return None 

298 

299 @staticmethod 

300 def _get_assembly_base_url_from_region(region: Literal["eu"] | None) -> str: 

301 """ 

302 Get the base URL for the AssemblyAI API 

303 if region == "eu", return "https://api.eu.assemblyai.com" 

304 else return "https://api.assemblyai.com" 

305 """ 

306 if region == "eu": 306 ↛ 307line 306 didn't jump to line 307 because the condition on line 306 was never true

307 return "https://api.eu.assemblyai.com" 

308 return "https://api.assemblyai.com"