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
« 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
9import httpx
10from typing_extensions import TypedDict
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)
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
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 """
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 """
51 self.max_polling_attempts = ASSEMBLY_AI_MAX_POLLING_ATTEMPTS
52 """
53 The maximum number of polling attempts for the AssemblyAI API.
54 """
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 )
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
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
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
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 )
134 passthrough_logging_payload: Final[PassthroughStandardLoggingPayload | None] = kwargs.get(
135 "passthrough_logging_payload"
136 )
138 verbose_proxy_logger.debug(
139 "standard_passthrough_logging_object %s",
140 json.dumps(passthrough_logging_payload, indent=4),
141 )
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
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 )
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)
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
174 Args:
175 response_body (dict): Response containing the transcript ID
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 )
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 }
201 response: Final = httpx.get(url, headers=headers)
202 response.raise_for_status()
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
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
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
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
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
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
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"
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
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"