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

349 statements  

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

1import asyncio 

2import json 

3import math 

4import tempfile 

5from collections.abc import Awaitable, Callable, Mapping 

6from dataclasses import dataclass 

7from datetime import datetime 

8from email.utils import parsedate_to_datetime 

9from functools import lru_cache, partial 

10from pathlib import Path 

11from types import MappingProxyType 

12from typing import IO, Final, Protocol, TypeAlias 

13from urllib.parse import quote 

14 

15import httpx 

16import soundfile 

17from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError 

18from typing_extensions import ReadOnly, TypedDict 

19 

20import litellm 

21from litellm._logging import verbose_proxy_logger 

22from litellm.constants import ( 

23 TRANSCRIBE_JOB_MAX_POLLING_ATTEMPTS, 

24 TRANSCRIBE_JOB_POLLING_INTERVAL_SECONDS, 

25 TRANSCRIBE_MAX_MEDIA_BYTES, 

26 TRANSCRIBE_MAX_MEDIA_DURATION_SECONDS, 

27 TRANSCRIBE_MEASURABLE_MEDIA_FORMATS, 

28 TRANSCRIBE_MEDIA_DOWNLOAD_CONCURRENCY, 

29 TRANSCRIBE_MEDIA_FETCH_ATTEMPTS, 

30 TRANSCRIBE_MEDIA_LAST_MODIFIED_TOLERANCE_SECONDS, 

31) 

32from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix 

33from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

34from litellm.litellm_core_utils.litellm_logging import ( 

35 get_standard_logging_object_payload, 

36) 

37from litellm.llms.custom_httpx.http_handler import get_async_httpx_client 

38from litellm.proxy._types import ( 

39 PassThroughEndpointLoggingResultValues, 

40 PassThroughEndpointLoggingTypedDict, 

41 UserAPIKeyAuth, 

42) 

43from litellm.proxy.common_utils.resource_ownership import ( 

44 get_primary_resource_owner_scope, 

45 is_proxy_admin, 

46 user_can_access_resource_owner, 

47) 

48from litellm.types.llms.custom_http import httpxSpecialProvider 

49from litellm.types.utils import StandardPassThroughResponseObject 

50 

51TRANSCRIBE_TARGET_PREFIX: Final = "Transcribe" 

52TRANSCRIBE_CUSTOM_LLM_PROVIDER: Final = "transcribe" 

53TRANSCRIBE_PRICED_OPERATION: Final = "StartTranscriptionJob" 

54TRANSCRIBE_PRICED_MODEL: Final = f"{TRANSCRIBE_CUSTOM_LLM_PROVIDER}/{TRANSCRIBE_PRICED_OPERATION}" 

55TRANSCRIBE_UNPRICED_OPERATIONS: Final = frozenset( 

56 {"StartCallAnalyticsJob", "StartMedicalScribeJob", "StartMedicalTranscriptionJob"} 

57) 

58TRANSCRIBE_SURCHARGE_MEMBERS: Final = ("ContentRedaction", "ToxicityDetection") 

59TRANSCRIBE_TERMINAL_JOB_STATUSES: Final = frozenset({"COMPLETED", "FAILED"}) 

60TRANSCRIBE_MISSING_JOB_ERRORS: Final = frozenset({"BadRequestException", "NotFoundException"}) 

61TRANSCRIBE_OWNER_TAG: Final = "litellm-owner" 

62TRANSCRIBE_OWNED_JOB_OPERATIONS: Final = frozenset({"GetTranscriptionJob", "DeleteTranscriptionJob"}) 

63TRANSCRIBE_MEDIA_BUCKETS_SETTING: Final = "transcribe_media_buckets" 

64TRANSCRIBE_ROLE_MEMBERS: Final = ("DataAccessRoleArn", "JobExecutionSettings") 

65TRANSCRIBE_MEDIA_URI_MEMBERS: Final = ("MediaFileUri", "RedactedMediaFileUri") 

66 

67JobLookup: TypeAlias = Callable[[str], Awaitable[Mapping[str, object]]] # mutable-ok: Callable parameter syntax 

68MediaDurationProbe: TypeAlias = Callable[[str, float], Awaitable[float | None]] # mutable-ok: Callable parameter syntax 

69 

70 

71class GetTranscriptionJobRequest(TypedDict): 

72 TranscriptionJobName: ReadOnly[str] 

73 

74 

75class _MediaRef(BaseModel): 

76 model_config = ConfigDict(frozen=True) 

77 MediaFileUri: str | None = None 

78 

79 

80class _JobTag(BaseModel): 

81 model_config = ConfigDict(frozen=True) 

82 Key: str | None = None 

83 Value: str | None = None 

84 

85 

86class TranscriptionJobRecord(BaseModel): 

87 model_config = ConfigDict(frozen=True) 

88 TranscriptionJobStatus: str | None = None 

89 CreationTime: float | None = None 

90 Media: _MediaRef | None = None 

91 Tags: tuple[_JobTag, ...] = () 

92 

93 

94class _TranscriptionJobResponse(BaseModel): 

95 model_config = ConfigDict(frozen=True) 

96 TranscriptionJob: TranscriptionJobRecord | None = None 

97 

98 

99@dataclass(frozen=True, slots=True) 

100class MissingJob: 

101 """Transcribe no longer knows the job, so polling it again can never reach a terminal status.""" 

102 

103 

104StartedJob: TypeAlias = TranscriptionJobRecord | None 

105JobPricer: TypeAlias = Callable[[str, str, float, StartedJob], Awaitable[float]] # mutable-ok: Callable params 

106 

107 

108class _PricedCostMapEntry(BaseModel): 

109 model_config = ConfigDict(frozen=True, strict=True) 

110 input_cost_per_second: float 

111 

112 

113_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object]) 

114_JSON_OBJECTS: Final = TypeAdapter(tuple[Mapping[str, object], ...]) 

115_BUCKET_NAMES: Final = TypeAdapter(frozenset[str]) 

116 

117 

118@dataclass(frozen=True, slots=True) 

119class TranscribeRefusal: 

120 status_code: int 

121 detail: str 

122 

123 

124class PassThroughLogDispatch(Protocol): 

125 def __call__( 125 ↛ exitline 125 didn't return from function '__call__' because

126 self, 

127 *, 

128 logging_obj: LiteLLMLoggingObj, 

129 standard_logging_response_object: PassThroughEndpointLoggingResultValues | None, 

130 result: str, 

131 start_time: datetime, 

132 end_time: datetime, 

133 cache_hit: bool, 

134 **kwargs: object, # kwargs-ok: mirrors the shared pass-through logging dispatch signature 

135 ) -> Awaitable[None]: ... 

136 

137 

138@lru_cache(maxsize=1) 

139def transcribe_supported_operations() -> frozenset[str]: 

140 """ 

141 Operation names of the Amazon Transcribe JSON 1.1 API, read from the botocore 

142 service model so the allowlist tracks the installed SDK instead of a hand-typed copy. 

143 """ 

144 from botocore.session import get_session 

145 

146 return frozenset(get_session().get_service_model("transcribe").operation_names) 

147 

148 

149def transcribe_cost_per_second() -> float | None: 

150 try: 

151 return _PricedCostMapEntry.model_validate(litellm.model_cost.get(TRANSCRIBE_PRICED_MODEL)).input_cost_per_second 

152 except ValidationError: 

153 return None 

154 

155 

156def transcribe_unpriceable_request_reason( 

157 operation: str, 

158 request_body: Mapping[str, object], 

159 cost_per_second: float | None, 

160) -> str | None: 

161 if operation in TRANSCRIBE_UNPRICED_OPERATIONS: 

162 return ( 

163 f"{operation} is billed per second of audio at a rate LiteLLM does not price yet, so it cannot be" 

164 f" submitted through this route; only {TRANSCRIBE_PRICED_OPERATION} is priced and budgeted" 

165 ) 

166 if operation != TRANSCRIBE_PRICED_OPERATION: 

167 return None 

168 if cost_per_second is None: 

169 return ( 

170 f"{TRANSCRIBE_PRICED_MODEL} has no input_cost_per_second in the LiteLLM model cost map, so billable" 

171 " transcription jobs cannot be submitted through this route" 

172 ) 

173 surcharges: Final = tuple(m for m in TRANSCRIBE_SURCHARGE_MEMBERS if m in request_body) + tuple( 

174 _custom_language_model_members(request_body) 

175 ) 

176 if surcharges: 

177 return ( 

178 f"{TRANSCRIBE_PRICED_OPERATION} with {', '.join(surcharges)} adds a per-second surcharge LiteLLM does not" 

179 " price yet; remove it to submit the job through this route" 

180 ) 

181 if requested_media_format(request_body) not in TRANSCRIBE_MEASURABLE_MEDIA_FORMATS: 

182 return ( 

183 "LiteLLM bills a transcription job by reading the length of the media file, which it can only do for" 

184 f" {', '.join(sorted(TRANSCRIBE_MEASURABLE_MEDIA_FORMATS))}; set MediaFormat to one of those or point" 

185 " Media.MediaFileUri at a file with that extension" 

186 ) 

187 return None 

188 

189 

190def _custom_language_model_members(request_body: Mapping[str, object]) -> tuple[str, ...]: 

191 model_settings: Final = request_body.get("ModelSettings") 

192 language_id_settings: Final = request_body.get("LanguageIdSettings") 

193 from_model_settings: Final = ( 

194 ("ModelSettings.LanguageModelName",) 

195 if isinstance(model_settings, Mapping) and "LanguageModelName" in model_settings 

196 else () 

197 ) 

198 from_language_id: Final = ( 

199 tuple( 

200 f"LanguageIdSettings.{language}.LanguageModelName" 

201 for language, settings in _JSON_OBJECT.validate_python(language_id_settings).items() 

202 if isinstance(settings, Mapping) and "LanguageModelName" in settings 

203 ) 

204 if isinstance(language_id_settings, Mapping) 

205 else () 

206 ) 

207 return from_model_settings + from_language_id 

208 

209 

210def requested_media_format(request_body: Mapping[str, object]) -> str | None: 

211 media_format: Final = request_body.get("MediaFormat") 

212 if isinstance(media_format, str): 

213 return media_format.lower() 

214 media: Final = request_body.get("Media") 

215 media_uri: Final = _JSON_OBJECT.validate_python(media).get("MediaFileUri") if isinstance(media, Mapping) else None 

216 if not isinstance(media_uri, str): 

217 return None 

218 path: Final = httpx.URL(media_uri).path if "://" in media_uri else media_uri 

219 _, dot, suffix = path.rpartition(".") 

220 return suffix.lower() if dot else None 

221 

222 

223def transcribe_admin_only_refusal(operation: str, user_api_key_dict: UserAPIKeyAuth) -> TranscribeRefusal | None: 

224 if ( 

225 operation == TRANSCRIBE_PRICED_OPERATION 

226 or operation in TRANSCRIBE_OWNED_JOB_OPERATIONS 

227 or is_proxy_admin(user_api_key_dict) 

228 ): 

229 return None 

230 return TranscribeRefusal( 

231 403, 

232 f"{operation} reaches every Amazon Transcribe resource in the AWS account, so only a proxy admin may call it;" 

233 f" other keys may {TRANSCRIBE_PRICED_OPERATION} and {' or '.join(sorted(TRANSCRIBE_OWNED_JOB_OPERATIONS))}" 

234 " for the jobs they started", 

235 ) 

236 

237 

238def transcribe_media_buckets(general_settings: Mapping[str, object]) -> frozenset[str] | None: 

239 try: 

240 return _BUCKET_NAMES.validate_python(general_settings.get(TRANSCRIBE_MEDIA_BUCKETS_SETTING)) 

241 except ValidationError: 

242 return None 

243 

244 

245def s3_bucket_name(uri: object) -> str | None: 

246 if not isinstance(uri, str) or not uri.startswith("s3://"): 

247 return None 

248 bucket, _, _ = uri.removeprefix("s3://").partition("/") 

249 return bucket or None 

250 

251 

252def transcribe_storage_refusal( 

253 request_body: Mapping[str, object], 

254 allowed_buckets: frozenset[str] | None, 

255 user_api_key_dict: UserAPIKeyAuth, 

256) -> TranscribeRefusal | None: 

257 """ 

258 Transcribe reads the media and writes the transcript with the proxy's own AWS credentials, so a 

259 non-admin key may only point a job at buckets the operator listed; otherwise any object those 

260 credentials can reach could be transcribed and read back through the caller's own job. 

261 """ 

262 if is_proxy_admin(user_api_key_dict): 

263 return None 

264 if allowed_buckets is None: 

265 return TranscribeRefusal( 

266 403, 

267 f"general_settings.{TRANSCRIBE_MEDIA_BUCKETS_SETTING} is not a list of S3 bucket names, so only a proxy" 

268 f" admin may {TRANSCRIBE_PRICED_OPERATION}; list the buckets other keys may read media from and write" 

269 " transcripts to", 

270 ) 

271 roles: Final = tuple(m for m in TRANSCRIBE_ROLE_MEMBERS if m in request_body) 

272 if roles: 

273 return TranscribeRefusal( 

274 403, 

275 f"{', '.join(roles)} would run the job under a role other than the proxy's own AWS credentials, so" 

276 " only a proxy admin may set it", 

277 ) 

278 media: Final = request_body.get("Media") 

279 media_uris: Final = ( 

280 tuple((f"Media.{m}", s3_bucket_name(media.get(m))) for m in TRANSCRIBE_MEDIA_URI_MEMBERS if m in media) 

281 if isinstance(media, Mapping) 

282 else () 

283 ) 

284 output: Final = request_body.get("OutputBucketName") 

285 locations: Final = media_uris + ( 

286 (("OutputBucketName", output if isinstance(output, str) else None),) 

287 if "OutputBucketName" in request_body 

288 else () 

289 ) 

290 offending: Final = tuple(member for member, bucket in locations if bucket not in allowed_buckets) 

291 if offending: 

292 return TranscribeRefusal( 

293 403, 

294 f"{', '.join(offending)} must name one of the S3 buckets in general_settings." 

295 f"{TRANSCRIBE_MEDIA_BUCKETS_SETTING} ({', '.join(sorted(allowed_buckets))}), as s3://bucket/key for media", 

296 ) 

297 return None 

298 

299 

300def transcribe_owned_start_request( 

301 request_body: Mapping[str, object], user_api_key_dict: UserAPIKeyAuth 

302) -> dict[str, object] | TranscribeRefusal: 

303 owner: Final = get_primary_resource_owner_scope(user_api_key_dict) 

304 if owner is None: 

305 return TranscribeRefusal(400, "The calling key has no identity to record as the owner of the transcription job") 

306 try: 

307 tags: Final = _JSON_OBJECTS.validate_python(request_body.get("Tags", ())) 

308 except ValidationError: 

309 return TranscribeRefusal(400, "Tags must be a list of objects with Key and Value members") 

310 if any(tag.get("Key") == TRANSCRIBE_OWNER_TAG for tag in tags): 

311 return TranscribeRefusal( 

312 400, f"The {TRANSCRIBE_OWNER_TAG} tag is assigned by LiteLLM and cannot be supplied by the caller" 

313 ) 

314 owner_tag: Final = _JobTag(Key=TRANSCRIBE_OWNER_TAG, Value=owner).model_dump() 

315 return {**request_body, "Tags": (*tags, owner_tag)} # mutable-ok: json.dumps and the body state key take a dict 

316 

317 

318async def transcribe_job_access_refusal( 

319 job_name: object, user_api_key_dict: UserAPIKeyAuth, get_job: JobLookup 

320) -> TranscribeRefusal | None: 

321 if is_proxy_admin(user_api_key_dict): 

322 return None 

323 if not isinstance(job_name, str): 

324 return TranscribeRefusal(400, "TranscriptionJobName must be a string") 

325 not_found: Final = TranscribeRefusal( 

326 404, f"No transcription job named {job_name} was started through this proxy by the calling key" 

327 ) 

328 try: 

329 job: Final = _TranscriptionJobResponse.model_validate(await get_job(job_name)).TranscriptionJob 

330 except Exception as e: # noqa: BLE001 # a job that cannot be read cannot be shown to belong to the caller 

331 verbose_proxy_logger.warning("Looking up Transcribe job %s for an ownership check failed: %s", job_name, e) 

332 return not_found 

333 owner: Final = ( 

334 next((tag.Value for tag in job.Tags if tag.Key == TRANSCRIBE_OWNER_TAG), None) if job is not None else None 

335 ) 

336 return None if user_can_access_resource_owner(owner, user_api_key_dict) else not_found 

337 

338 

339def transcription_job_cost(audio_seconds: float, cost_per_second: float) -> float: 

340 return math.ceil(audio_seconds) * cost_per_second 

341 

342 

343def transcribe_max_job_cost(cost_per_second: float) -> float: 

344 return transcription_job_cost(TRANSCRIBE_MAX_MEDIA_DURATION_SECONDS, cost_per_second) 

345 

346 

347def started_transcription_job(response_body: Mapping[str, object] | None) -> TranscriptionJobRecord | None: 

348 try: 

349 return _TranscriptionJobResponse.model_validate(response_body).TranscriptionJob 

350 except ValidationError: 

351 return None 

352 

353 

354def aws_error_type(response: httpx.Response) -> str | None: 

355 try: 

356 error_type: Final = _JSON_OBJECT.validate_python(response.json()).get("__type") 

357 except (ValueError, ValidationError): 

358 return None 

359 return error_type.rsplit("#", 1)[-1] if isinstance(error_type, str) else None 

360 

361 

362async def _poll_transcription_job(job_name: str, get_job: JobLookup) -> TranscriptionJobRecord | MissingJob | None: 

363 try: 

364 job: Final = _TranscriptionJobResponse.model_validate(await get_job(job_name)).TranscriptionJob 

365 except httpx.HTTPStatusError as e: 

366 if aws_error_type(e.response) in TRANSCRIBE_MISSING_JOB_ERRORS: 

367 verbose_proxy_logger.warning( 

368 "Transcribe job %s no longer exists, pricing the media it was started with", job_name 

369 ) 

370 return MissingJob() 

371 verbose_proxy_logger.warning("Polling Transcribe job %s failed, retrying: %s", job_name, e) 

372 return None 

373 except Exception as e: # noqa: BLE001 # a failed poll is retried on the next tick instead of ending pricing 

374 verbose_proxy_logger.warning("Polling Transcribe job %s failed, retrying: %s", job_name, e) 

375 return None 

376 return job if job is not None and job.TranscriptionJobStatus in TRANSCRIBE_TERMINAL_JOB_STATUSES else None 

377 

378 

379async def await_transcription_job( 

380 job_name: str, 

381 get_job: JobLookup, 

382 sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, 

383 max_attempts: int = TRANSCRIBE_JOB_MAX_POLLING_ATTEMPTS, 

384) -> TranscriptionJobRecord | MissingJob | None: 

385 for _ in range(max_attempts): 

386 job = await _poll_transcription_job(job_name, get_job) 

387 if job is not None: 

388 return job 

389 await sleep(TRANSCRIBE_JOB_POLLING_INTERVAL_SECONDS) 

390 return None 

391 

392 

393async def measure_media_seconds( 

394 media_uri: str, 

395 job_created_at: float, 

396 media_seconds: MediaDurationProbe, 

397 sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, 

398 attempts: int = TRANSCRIBE_MEDIA_FETCH_ATTEMPTS, 

399) -> float | None: 

400 for attempt in range(1, attempts + 1): 

401 try: 

402 return await media_seconds(media_uri, job_created_at) 

403 except Exception as e: # noqa: BLE001 # the media is retried, then charged at the maximum if still unreadable 

404 verbose_proxy_logger.warning("Measuring Transcribe media %s failed (attempt %d): %s", media_uri, attempt, e) 

405 if attempt < attempts: 

406 await sleep(TRANSCRIBE_JOB_POLLING_INTERVAL_SECONDS) 

407 return None 

408 

409 

410async def price_transcription_job( 

411 job_name: str, 

412 cost_per_second: float, 

413 get_job: JobLookup, 

414 media_seconds: MediaDurationProbe, 

415 sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, 

416 max_attempts: int = TRANSCRIBE_JOB_MAX_POLLING_ATTEMPTS, 

417 started_job: TranscriptionJobRecord | None = None, 

418) -> float: 

419 """ 

420 Amazon Transcribe bills every second of the media file, silence included, and reports no 

421 duration itself, so the job is polled to completion and the media it transcribed is measured. 

422 The measurement only counts when the object has not been rewritten since the job was created, 

423 which is what ties it to the bytes Transcribe read. A job deleted before it is polled is 

424 measured from the media named in its StartTranscriptionJob response. Anything that stops the 

425 duration from being read is charged as the longest media AWS accepts. 

426 """ 

427 outcome: Final = await await_transcription_job(job_name, get_job, sleep=sleep, max_attempts=max_attempts) 

428 if outcome is None: 

429 verbose_proxy_logger.warning("Transcribe job %s did not finish while polling, charging maximum", job_name) 

430 return transcribe_max_job_cost(cost_per_second) 

431 if isinstance(outcome, TranscriptionJobRecord) and outcome.TranscriptionJobStatus == "FAILED": 

432 return 0.0 

433 job: Final = outcome if isinstance(outcome, TranscriptionJobRecord) else started_job 

434 media_uri: Final = job.Media.MediaFileUri if job is not None and job.Media is not None else None 

435 if job is None or media_uri is None or job.CreationTime is None: 

436 return transcribe_max_job_cost(cost_per_second) 

437 audio_seconds: Final = await measure_media_seconds(media_uri, job.CreationTime, media_seconds, sleep=sleep) 

438 if audio_seconds is None: 

439 return transcribe_max_job_cost(cost_per_second) 

440 return transcription_job_cost(audio_seconds, cost_per_second) 

441 

442 

443def _as_json_object(response: httpx.Response) -> Mapping[str, object]: 

444 return _JSON_OBJECT.validate_python(response.raise_for_status().json()) 

445 

446 

447def transcribe_job_lookup(aws_region_name: str) -> JobLookup: 

448 from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing, sign_aws_json_post 

449 

450 url: Final = f"https://transcribe.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}/" 

451 headers: Final = MappingProxyType( 

452 { 

453 "Content-Type": "application/x-amz-json-1.1", 

454 "X-Amz-Target": f"{TRANSCRIBE_TARGET_PREFIX}.GetTranscriptionJob", 

455 } 

456 ) 

457 

458 async def get_job(job_name: str) -> Mapping[str, object]: 

459 body: Final[GetTranscriptionJobRequest] = {"TranscriptionJobName": job_name} 

460 payload: Final = json.dumps(body) 

461 prepped: Final = await run_aws_signing( 

462 sign_aws_json_post, 

463 get_credentials=partial(BaseAWSLLM().get_credentials, aws_region_name=aws_region_name), 

464 service_name="transcribe", 

465 aws_region_name=aws_region_name, 

466 url=url, 

467 body=payload, 

468 headers=headers, 

469 ) 

470 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.PassThroughEndpoint) 

471 signed_headers: Final = dict(prepped.headers.items()) # mutable-ok: AsyncHTTPHandler.post takes a dict 

472 return _as_json_object(await client.post(str(prepped.url), data=payload, headers=signed_headers)) 

473 

474 return get_job 

475 

476 

477def s3_media_url(media_uri: str, aws_region_name: str) -> str | None: 

478 """ 

479 Transcribe accepts media as s3://bucket/key or as an https S3 URL; the bucket is required to 

480 live in the job's region, so the s3 form maps onto that region's endpoint. Buckets with dots in 

481 their name use the path-style form because they cannot match the virtual-hosted wildcard 

482 certificate. The proxy's AWS signature is only ever sent to that partition's own hosts. 

483 """ 

484 dns_suffix: Final = get_aws_dns_suffix(aws_region_name) 

485 if not media_uri.startswith("s3://"): 

486 url: Final = httpx.URL(media_uri) 

487 return media_uri if url.scheme == "https" and url.host.endswith(f".{dns_suffix}") else None 

488 bucket, _, key = media_uri.removeprefix("s3://").partition("/") 

489 if "." in bucket: 

490 return f"https://s3.{aws_region_name}.{dns_suffix}/{bucket}/{quote(key)}" 

491 return f"https://{bucket}.s3.{aws_region_name}.{dns_suffix}/{quote(key)}" 

492 

493 

494def media_predates_job(headers: Mapping[str, str], job_created_at: float) -> bool: 

495 try: 

496 modified_at: Final = parsedate_to_datetime(headers["last-modified"]).timestamp() 

497 except (KeyError, TypeError, ValueError): 

498 return False 

499 return modified_at <= job_created_at + TRANSCRIBE_MEDIA_LAST_MODIFIED_TOLERANCE_SECONDS 

500 

501 

502async def write_media_within_limit(response: httpx.Response, media_file: IO[bytes], max_bytes: int) -> bool: 

503 if int(response.headers.get("content-length", "0")) > max_bytes: 

504 return False 

505 async for chunk in response.aiter_bytes(): 

506 _ = media_file.write(chunk) 

507 if media_file.tell() > max_bytes: 

508 return False 

509 return True 

510 

511 

512def media_file_seconds(path: Path) -> float | None: 

513 try: 

514 with soundfile.SoundFile(str(path)) as audio: 

515 return len(audio) / audio.samplerate 

516 except (RuntimeError, ValueError, OSError) as e: 

517 verbose_proxy_logger.warning("Transcribe media could not be decoded for its duration: %s", e) 

518 return None 

519 

520 

521def transcribe_media_duration_probe(aws_region_name: str, download_slots: asyncio.Semaphore) -> MediaDurationProbe: 

522 from botocore.auth import S3SigV4Auth 

523 from botocore.awsrequest import AWSRequest 

524 

525 from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing 

526 

527 def sign_s3_get(url: str) -> dict[str, str]: # mutable-ok: httpx request headers take a dict 

528 aws_request: Final = AWSRequest(method="GET", url=url) 

529 credentials: Final = BaseAWSLLM().get_credentials(aws_region_name=aws_region_name) 

530 S3SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request) 

531 return dict(aws_request.prepare().headers.items()) # mutable-ok: httpx request headers take a dict 

532 

533 async def media_seconds(media_uri: str, job_created_at: float) -> float | None: 

534 url: Final = s3_media_url(media_uri, aws_region_name) 

535 if url is None: 

536 return None 

537 headers: Final = await run_aws_signing(sign_s3_get, url) 

538 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.PassThroughEndpoint).client 

539 async with download_slots: 

540 with tempfile.NamedTemporaryFile() as media_file: 

541 async with client.stream("GET", url, headers=headers) as response: 

542 _ = response.raise_for_status() 

543 if not media_predates_job(response.headers, job_created_at): 

544 verbose_proxy_logger.warning( 

545 "Transcribe media %s was rewritten after the job was created, charging maximum", media_uri 

546 ) 

547 return None 

548 if not await write_media_within_limit(response, media_file, TRANSCRIBE_MAX_MEDIA_BYTES): 

549 verbose_proxy_logger.warning( 

550 "Transcribe media %s exceeds the size cap, charging maximum", media_uri 

551 ) 

552 return None 

553 media_file.flush() 

554 return await asyncio.to_thread(media_file_seconds, Path(media_file.name)) 

555 

556 return media_seconds 

557 

558 

559async def price_transcription_job_live( 

560 job_name: str, 

561 aws_region_name: str, 

562 cost_per_second: float, 

563 started_job: TranscriptionJobRecord | None, 

564 download_slots: asyncio.Semaphore, 

565) -> float: 

566 try: 

567 return await price_transcription_job( 

568 job_name, 

569 cost_per_second, 

570 get_job=transcribe_job_lookup(aws_region_name), 

571 media_seconds=transcribe_media_duration_probe(aws_region_name, download_slots), 

572 started_job=started_job, 

573 ) 

574 except Exception as e: # noqa: BLE001 # an unreadable job must still be charged, so fail closed at the maximum 

575 verbose_proxy_logger.exception("Pricing Transcribe job %s failed, charging maximum: %s", job_name, e) 

576 return transcribe_max_job_cost(cost_per_second) 

577 

578 

579class TranscribePassthroughLoggingHandler: 

580 def __init__(self, job_pricer: JobPricer | None = None) -> None: 

581 self._job_pricer: Final = ( 

582 job_pricer 

583 if job_pricer is not None 

584 else partial( 

585 price_transcription_job_live, 

586 download_slots=asyncio.Semaphore(TRANSCRIBE_MEDIA_DOWNLOAD_CONCURRENCY), 

587 ) 

588 ) 

589 self._pricing_tasks: Final[set[asyncio.Task[None]]] = set() # mutable-ok: asyncio holds tasks weakly 

590 

591 @staticmethod 

592 def _operation_from_response(httpx_response: httpx.Response) -> str: 

593 headers: Final[Mapping[str, str]] = httpx_response.request.headers 

594 target: Final = headers.get("x-amz-target", "") 

595 return target.split(".")[-1] 

596 

597 @staticmethod 

598 def is_priced_job_start(httpx_response: httpx.Response) -> bool: 

599 return ( 

600 TranscribePassthroughLoggingHandler._operation_from_response(httpx_response) == TRANSCRIBE_PRICED_OPERATION 

601 ) 

602 

603 def schedule_priced_job_logging( 

604 self, 

605 httpx_response: httpx.Response, 

606 response_body: Mapping[str, object] | None, 

607 logging_obj: LiteLLMLoggingObj, 

608 url_route: str, 

609 result: str, 

610 start_time: datetime, 

611 end_time: datetime, 

612 cache_hit: bool, 

613 request_body: Mapping[str, object], 

614 log: PassThroughLogDispatch, 

615 **kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler 

616 ) -> asyncio.Task[None]: 

617 task: Final = asyncio.create_task( 

618 self._price_then_log( 

619 httpx_response=httpx_response, 

620 started_job=started_transcription_job(response_body), 

621 logging_obj=logging_obj, 

622 url_route=url_route, 

623 result=result, 

624 start_time=start_time, 

625 end_time=end_time, 

626 cache_hit=cache_hit, 

627 request_body=request_body, 

628 log=log, 

629 **kwargs, 

630 ) 

631 ) 

632 self._pricing_tasks.add(task) 

633 task.add_done_callback(self._pricing_tasks.discard) 

634 return task 

635 

636 async def _price_then_log( 

637 self, 

638 httpx_response: httpx.Response, 

639 started_job: TranscriptionJobRecord | None, 

640 logging_obj: LiteLLMLoggingObj, 

641 url_route: str, 

642 result: str, 

643 start_time: datetime, 

644 end_time: datetime, 

645 cache_hit: bool, 

646 request_body: Mapping[str, object], 

647 log: PassThroughLogDispatch, 

648 **kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler 

649 ) -> None: 

650 cost_per_second: Final = transcribe_cost_per_second() 

651 if cost_per_second is None: 

652 verbose_proxy_logger.error("%s left the model cost map, spend not recorded", TRANSCRIBE_PRICED_MODEL) 

653 return 

654 job_name: Final = request_body.get("TranscriptionJobName") 

655 aws_region_name: Final = httpx_response.request.url.host.split(".")[1] 

656 response_cost: Final = await self._job_pricer( 

657 job_name if isinstance(job_name, str) else "", 

658 aws_region_name, 

659 cost_per_second, 

660 started_job, 

661 ) 

662 payload: Final = self.transcribe_passthrough_handler( 

663 httpx_response=httpx_response, 

664 logging_obj=logging_obj, 

665 url_route=url_route, 

666 result=result, 

667 start_time=start_time, 

668 end_time=end_time, 

669 cache_hit=cache_hit, 

670 request_body=request_body, 

671 response_cost=response_cost, 

672 **kwargs, 

673 ) 

674 await log( 

675 logging_obj=logging_obj, 

676 standard_logging_response_object=payload["result"], 

677 result=result, 

678 start_time=start_time, 

679 end_time=end_time, 

680 cache_hit=cache_hit, 

681 **payload["kwargs"], 

682 ) 

683 

684 @staticmethod 

685 def transcribe_passthrough_handler( 

686 httpx_response: httpx.Response, 

687 logging_obj: LiteLLMLoggingObj, 

688 url_route: str, 

689 result: str, 

690 start_time: datetime, 

691 end_time: datetime, 

692 cache_hit: bool, 

693 request_body: Mapping[str, object], 

694 response_cost: float = 0.0, 

695 **kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler 

696 ) -> PassThroughEndpointLoggingTypedDict: 

697 try: 

698 operation: Final = TranscribePassthroughLoggingHandler._operation_from_response(httpx_response) 

699 model_name: Final = f"{TRANSCRIBE_CUSTOM_LLM_PROVIDER}/{operation}" 

700 

701 updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict 

702 **kwargs, 

703 "model": model_name, 

704 "custom_llm_provider": TRANSCRIBE_CUSTOM_LLM_PROVIDER, 

705 "response_cost": response_cost, 

706 } 

707 logging_obj.model_call_details.update( 

708 model=model_name, 

709 custom_llm_provider=TRANSCRIBE_CUSTOM_LLM_PROVIDER, 

710 response_cost=response_cost, 

711 ) 

712 

713 standard_logging_object: Final = get_standard_logging_object_payload( 

714 kwargs=updated_kwargs, 

715 init_response_obj=StandardPassThroughResponseObject(response=result), 

716 start_time=start_time, 

717 end_time=end_time, 

718 logging_obj=logging_obj, 

719 status="success", 

720 ) 

721 

722 handler_payload: Final[PassThroughEndpointLoggingTypedDict] = { 

723 "result": StandardPassThroughResponseObject(response=result), 

724 "kwargs": {**updated_kwargs, "standard_logging_object": standard_logging_object}, 

725 } 

726 except Exception as e: # noqa: BLE001 # logging must never fail the forwarded request 

727 verbose_proxy_logger.exception("Error in Amazon Transcribe passthrough logging handler: %s", e) 

728 fallback_payload: Final[PassThroughEndpointLoggingTypedDict] = { 

729 "result": StandardPassThroughResponseObject(response=result), 

730 "kwargs": kwargs, 

731 } 

732 return fallback_payload 

733 return handler_payload