Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/container_endpoints/ownership.py: 29%

213 statements  

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

1import json 

2from collections.abc import Awaitable, Callable, Mapping, Sequence 

3from collections.abc import Set as AbstractSet 

4from typing import TYPE_CHECKING, Any, Final, TypeAlias 

5 

6from fastapi import HTTPException 

7from pydantic import BaseModel 

8 

9from litellm._logging import verbose_proxy_logger 

10from litellm.caching.in_memory_cache import InMemoryCache 

11from litellm.proxy._types import UserAPIKeyAuth 

12from litellm.proxy.common_utils.resource_ownership import ( 

13 get_primary_resource_owner_scope, 

14 get_resource_owner_scopes, 

15 is_proxy_admin, 

16 user_can_access_resource_owner, 

17) 

18from litellm.repositories.table_repositories import ManagedObjectRepository 

19from litellm.responses.utils import ResponsesAPIRequestUtils 

20 

21if TYPE_CHECKING: 21 ↛ 22line 21 didn't jump to line 22 because the condition on line 21 was never true

22 from prisma import models as prisma_models 

23 

24 from litellm.proxy.utils import PrismaClient 

25 

26 

27CONTAINER_OBJECT_PURPOSE: Final = "container" 

28 

29# 60s LRU/TTL cache absorbs every container access check before it reaches 

30# Prisma. ``_NEGATIVE_OWNER_SENTINEL`` lets us cache a true "untracked" 

31# answer so repeated misses also avoid the DB — ``InMemoryCache`` returns 

32# ``None`` indistinguishably for "miss" and "cached as None". 

33_NEGATIVE_OWNER_SENTINEL: Final = "__litellm_container_no_owner__" 

34_CONTAINER_OWNER_CACHE: Final = InMemoryCache(max_size_in_memory=10000, default_ttl=60) 

35 

36# Caches the stored ``unified_object_id`` (the encoded container ID 

37# captured at create time) so ``get_container_forwarding_params`` can 

38# recover the deployment ``model_id`` for native upstream IDs without 

39# re-hitting Prisma on every retrieve/delete. 

40_NEGATIVE_STORED_ID_SENTINEL: Final = "__litellm_container_no_stored_id__" 

41_CONTAINER_STORED_ID_CACHE: Final = InMemoryCache(max_size_in_memory=10000, default_ttl=60) 

42 

43# Per-caller-scope cache for ``GET /v1/containers`` list filtering. Without 

44# this, every list call issues a fresh ``find_many`` against 

45# ``litellm_managedobjecttable``. The cache key is the sorted owner-scope 

46# tuple — different keys for the same user share the same allow-set, but 

47# different users with different scopes get disjoint cache entries. 

48_ALLOWED_CONTAINER_IDS_CACHE: Final = InMemoryCache(max_size_in_memory=2048, default_ttl=60) 

49 

50DEFAULT_CONTAINER_LIST_LIMIT: Final = 20 

51OWNED_CONTAINER_LIST_PAGE_SIZE: Final = 100 

52OWNED_CONTAINER_LIST_MAX_PAGES: Final = 5 

53 

54FetchContainerListPage: TypeAlias = Callable[[str | None, int | None], Awaitable[object]] 

55 

56 

57def _allowed_container_ids_cache_key(owner_scopes: Sequence[str]) -> str: 

58 """JSON-encode the sorted scope list — using a separator like ``|`` 

59 would collide for any tenant whose user_id / team_id / org_id / 

60 api_key happens to contain the separator. JSON quoting escapes 

61 every separator that matters.""" 

62 return json.dumps(sorted(owner_scopes)) 

63 

64 

65def _container_model_object_id(original_container_id: str, custom_llm_provider: str) -> str: 

66 return f"{CONTAINER_OBJECT_PURPOSE}:{custom_llm_provider}:{original_container_id}" 

67 

68 

69def decode_container_id_for_ownership(container_id: str, custom_llm_provider: str) -> tuple[str, str]: 

70 decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) 

71 original_container_id: Final = decoded.get("response_id", container_id) 

72 decoded_provider: Final = decoded.get("custom_llm_provider") 

73 if decoded_provider and custom_llm_provider == "openai": 73 ↛ 74line 73 didn't jump to line 74 because the condition on line 73 was never true

74 custom_llm_provider = decoded_provider 

75 return original_container_id, custom_llm_provider 

76 

77 

78async def get_container_forwarding_params( 

79 container_id: str, original_container_id: str, custom_llm_provider: str 

80) -> dict[str, str]: 

81 params: Final = { 

82 "container_id": original_container_id, 

83 "custom_llm_provider": custom_llm_provider, 

84 } 

85 decoded: Final = ResponsesAPIRequestUtils._decode_container_id(container_id) 

86 model_id = decoded.get("model_id") 

87 if not (isinstance(model_id, str) and model_id): 87 ↛ 99line 87 didn't jump to line 99 because the condition on line 87 was always true

88 # Native upstream IDs (e.g. Azure ``cntr_<hex>``) carry no LiteLLM 

89 # routing payload, so decoding the user-supplied id yields no 

90 # ``model_id``. Recover it from the encoded ``unified_object_id`` 

91 # captured on the ownership row at create time — when the router 

92 # selected a specific deployment that ID embeds the model_id. 

93 stored_id: Final = await _get_stored_container_id(original_container_id, custom_llm_provider) 

94 if stored_id and stored_id != container_id: 94 ↛ 95line 94 didn't jump to line 95 because the condition on line 94 was never true

95 stored_decoded: Final = ResponsesAPIRequestUtils._decode_container_id(stored_id) 

96 stored_model_id: Final = stored_decoded.get("model_id") 

97 if isinstance(stored_model_id, str) and stored_model_id: 

98 model_id = stored_model_id 

99 if isinstance(model_id, str) and model_id: 99 ↛ 100line 99 didn't jump to line 100 because the condition on line 99 was never true

100 params["model_id"] = model_id 

101 return params 

102 

103 

104def _get_response_id(response: object) -> str | None: 

105 if response is None: 

106 return None 

107 if isinstance(response, dict): 

108 value = response.get("id") 

109 else: 

110 value = getattr(response, "id", None) 

111 return value if isinstance(value, str) else None 

112 

113 

114def _dump_response(response: Any) -> dict[str, object]: 

115 if isinstance(response, dict): 

116 return dict(response) 

117 if hasattr(response, "model_dump"): 

118 return response.model_dump() 

119 if hasattr(response, "dict"): 

120 return response.dict() 

121 return {"id": _get_response_id(response)} 

122 

123 

124async def _get_prisma_client() -> "PrismaClient | None": 

125 from litellm.proxy.proxy_server import prisma_client 

126 

127 return prisma_client 

128 

129 

130def _custom_llm_provider_from_responses_response( 

131 response: object, 

132 default: str = "openai", 

133) -> str: 

134 hidden_params: Mapping[str, object] = {} 

135 if isinstance(response, dict): 

136 hidden_params = response.get("_hidden_params") or {} 

137 else: 

138 hidden_params = getattr(response, "_hidden_params", None) or {} 

139 

140 provider: Final = hidden_params.get("custom_llm_provider") 

141 if isinstance(provider, str) and provider: 

142 return provider 

143 return default 

144 

145 

146async def record_container_owners_from_responses_response( 

147 response: object, 

148 user_api_key_dict: UserAPIKeyAuth, 

149 custom_llm_provider: str | None = None, 

150) -> None: 

151 """Track containers created implicitly by code interpreter in /v1/responses.""" 

152 container_ids: Final = ResponsesAPIRequestUtils.collect_container_ids_from_responses_response(response) 

153 if not container_ids: 

154 return 

155 

156 resolved_provider: Final = custom_llm_provider or _custom_llm_provider_from_responses_response(response) 

157 

158 for container_id in container_ids: 

159 try: 

160 await record_container_owner( 

161 response={"id": container_id, "object": "container"}, 

162 user_api_key_dict=user_api_key_dict, 

163 custom_llm_provider=resolved_provider, 

164 ) 

165 except Exception as e: 

166 # Per-container errors (including ``HTTPException`` from 

167 # conflicting/forbidden ownership rows) must not abort the 

168 # batch — other containers in the same response should still 

169 # get recorded so their follow-up file API calls don't 403. 

170 verbose_proxy_logger.exception( 

171 "Failed to record container ownership from responses output for container_id=%s: %s", 

172 container_id, 

173 e, 

174 ) 

175 

176 

177async def record_container_owner( 

178 response: object, 

179 user_api_key_dict: UserAPIKeyAuth, 

180 custom_llm_provider: str, 

181) -> object: 

182 container_id: Final = _get_response_id(response) 

183 if container_id is None: 

184 verbose_proxy_logger.warning("Skipping container ownership tracking because provider response has no id") 

185 return response 

186 owner: Final = get_primary_resource_owner_scope(user_api_key_dict) 

187 if owner is None: 

188 # Admins with identity (the common path: master-key auth populates 

189 # ``user_id`` + ``api_key``) flow through the normal record path 

190 # below so admin-created containers are still tracked. Truly 

191 # identity-less admins (no user_id / team_id / org_id / api_key / 

192 # token) can't be uniquely stamped on the row — stamping a 

193 # placeholder would collapse every such caller into a shared 

194 # owner, the cross-tenant primitive we explicitly avoid. 

195 raise HTTPException( 

196 status_code=403, 

197 detail="Unable to record container ownership: caller has no identity scope.", 

198 ) 

199 

200 original_container_id, resolved_provider = decode_container_id_for_ownership(container_id, custom_llm_provider) 

201 model_object_id: Final = _container_model_object_id(original_container_id, resolved_provider) 

202 file_object: Final = _dump_response(response) 

203 file_object["custom_llm_provider"] = resolved_provider 

204 file_object["provider_container_id"] = original_container_id 

205 # Prisma Python requires Json fields to be serialized as a JSON string. 

206 file_object_json: Final[str] = json.dumps(file_object) 

207 

208 prisma_client: Final = await _get_prisma_client() 

209 if prisma_client is None: 

210 verbose_proxy_logger.warning("Skipping container ownership tracking because prisma_client is None") 

211 return response 

212 

213 table: Final = ManagedObjectRepository(prisma_client).table 

214 existing: Final = await table.find_unique(where={"model_object_id": model_object_id}) 

215 if existing is not None: 

216 if getattr(existing, "file_purpose", None) != CONTAINER_OBJECT_PURPOSE: 

217 raise HTTPException(status_code=500, detail="Unable to track container") 

218 if not user_can_access_resource_owner(getattr(existing, "created_by", None), user_api_key_dict): 

219 raise HTTPException(status_code=403, detail="Forbidden") 

220 await table.update( 

221 where={"model_object_id": model_object_id}, 

222 data={ 

223 "unified_object_id": container_id, 

224 "file_object": file_object_json, 

225 "updated_by": owner, 

226 }, 

227 ) 

228 else: 

229 await table.create( 

230 data={ 

231 "unified_object_id": container_id, 

232 "model_object_id": model_object_id, 

233 "file_object": file_object_json, 

234 "file_purpose": CONTAINER_OBJECT_PURPOSE, 

235 "created_by": owner, 

236 "updated_by": owner, 

237 } 

238 ) 

239 

240 _CONTAINER_OWNER_CACHE.set_cache(model_object_id, owner) 

241 _CONTAINER_STORED_ID_CACHE.set_cache(model_object_id, container_id) 

242 # Drop the caller's own list-cache entry so the just-created container 

243 # shows up on their next ``GET /v1/containers``. Other callers with 

244 # disjoint scope tuples have their own entries; intersecting-scope 

245 # tuples self-correct on the 60s TTL. 

246 caller_scopes: Final = get_resource_owner_scopes(user_api_key_dict) 

247 if caller_scopes: 

248 _ALLOWED_CONTAINER_IDS_CACHE.delete_cache(_allowed_container_ids_cache_key(caller_scopes)) 

249 return response 

250 

251 

252async def _get_container_owner(original_container_id: str, custom_llm_provider: str) -> str | None: 

253 model_object_id: Final = _container_model_object_id(original_container_id, custom_llm_provider) 

254 

255 cached: Final = _CONTAINER_OWNER_CACHE.get_cache(model_object_id) 

256 if cached == _NEGATIVE_OWNER_SENTINEL: 

257 return None 

258 if cached is not None: 

259 return cached 

260 

261 prisma_client: Final = await _get_prisma_client() 

262 if prisma_client is None: 

263 return None 

264 

265 table: Final = ManagedObjectRepository(prisma_client).table 

266 row: Final[prisma_models.LiteLLM_ManagedObjectTable | None] = await table.find_first( 

267 where={ 

268 "model_object_id": model_object_id, 

269 "file_purpose": CONTAINER_OBJECT_PURPOSE, 

270 } 

271 ) 

272 owner: Final[str | None] = getattr(row, "created_by", None) if row is not None else None 

273 _CONTAINER_OWNER_CACHE.set_cache(model_object_id, owner if owner is not None else _NEGATIVE_OWNER_SENTINEL) 

274 stored_id: Final[str | None] = getattr(row, "unified_object_id", None) if row is not None else None 

275 _CONTAINER_STORED_ID_CACHE.set_cache( 

276 model_object_id, 

277 (stored_id if isinstance(stored_id, str) and stored_id else _NEGATIVE_STORED_ID_SENTINEL), 

278 ) 

279 return owner 

280 

281 

282async def _get_stored_container_id(original_container_id: str, custom_llm_provider: str) -> str | None: 

283 """Return the ``unified_object_id`` stored at create time, if any. 

284 

285 Used by :func:`get_container_forwarding_params` to recover the 

286 deployment ``model_id`` for native upstream container IDs: the stored 

287 value is the encoded form produced by ``encode_container_id_in_response`` 

288 when the router selected a specific deployment. 

289 """ 

290 model_object_id: Final = _container_model_object_id(original_container_id, custom_llm_provider) 

291 

292 cached: Final = _CONTAINER_STORED_ID_CACHE.get_cache(model_object_id) 

293 if cached == _NEGATIVE_STORED_ID_SENTINEL: 

294 return None 

295 if isinstance(cached, str) and cached: 295 ↛ 296line 295 didn't jump to line 296 because the condition on line 295 was never true

296 return cached 

297 

298 prisma_client: Final = await _get_prisma_client() 

299 if prisma_client is None: 299 ↛ 300line 299 didn't jump to line 300 because the condition on line 299 was never true

300 return None 

301 

302 table: Final = ManagedObjectRepository(prisma_client).table 

303 row: Final[prisma_models.LiteLLM_ManagedObjectTable | None] = await table.find_first( 

304 where={ 

305 "model_object_id": model_object_id, 

306 "file_purpose": CONTAINER_OBJECT_PURPOSE, 

307 } 

308 ) 

309 stored_id: Final[str | None] = getattr(row, "unified_object_id", None) if row is not None else None 

310 _CONTAINER_STORED_ID_CACHE.set_cache( 

311 model_object_id, 

312 (stored_id if isinstance(stored_id, str) and stored_id else _NEGATIVE_STORED_ID_SENTINEL), 

313 ) 

314 return stored_id if isinstance(stored_id, str) and stored_id else None 

315 

316 

317async def assert_user_can_access_container( 

318 container_id: str, 

319 user_api_key_dict: UserAPIKeyAuth, 

320 custom_llm_provider: str, 

321) -> tuple[str, str]: 

322 original_container_id, resolved_provider = decode_container_id_for_ownership(container_id, custom_llm_provider) 

323 

324 if is_proxy_admin(user_api_key_dict): 324 ↛ 330line 324 didn't jump to line 330 because the condition on line 324 was always true

325 return original_container_id, resolved_provider 

326 

327 # Untracked rows (no ownership) are admin-only. Pre-isolation rows 

328 # that pre-date this enforcement need an admin to either re-create 

329 # via the now-tracked flow or assign ``created_by`` on the row. 

330 owner: Final = await _get_container_owner(original_container_id, resolved_provider) 

331 if not user_can_access_resource_owner(owner, user_api_key_dict): 

332 raise HTTPException(status_code=403, detail="Forbidden") 

333 

334 return original_container_id, resolved_provider 

335 

336 

337def _get_container_list_data(response: object) -> Sequence[object] | None: 

338 if response is None: 

339 return None 

340 if isinstance(response, dict): 

341 data = response.get("data") 

342 else: 

343 data = getattr(response, "data", None) 

344 return data if isinstance(data, list) else None 

345 

346 

347def _get_has_more(response: object) -> bool: 

348 if isinstance(response, dict): 

349 return response.get("has_more") is True 

350 return getattr(response, "has_more", None) is True 

351 

352 

353def _with_container_list_page(response: object, data: Sequence[object], has_more: bool) -> object: 

354 page: Final = { 

355 "data": list(data), 

356 "first_id": _get_response_id(data[0]) if data else None, 

357 "last_id": _get_response_id(data[-1]) if data else None, 

358 "has_more": has_more, 

359 } 

360 if isinstance(response, dict): 

361 return {**response, **page} 

362 if isinstance(response, BaseModel): 

363 return response.model_copy(update=page) 

364 return response 

365 

366 

367async def _get_allowed_container_ids( 

368 user_api_key_dict: UserAPIKeyAuth, 

369) -> AbstractSet[str]: 

370 owner_scopes: Final = get_resource_owner_scopes(user_api_key_dict) 

371 if not owner_scopes: 

372 return frozenset() 

373 

374 cache_key: Final = _allowed_container_ids_cache_key(owner_scopes) 

375 cached: Final = _ALLOWED_CONTAINER_IDS_CACHE.get_cache(cache_key) 

376 if cached is not None: 

377 return frozenset(cached) 

378 

379 prisma_client: Final = await _get_prisma_client() 

380 if prisma_client is None: 

381 return frozenset() 

382 

383 table: Final = ManagedObjectRepository(prisma_client).table 

384 rows: Final[Sequence[prisma_models.LiteLLM_ManagedObjectTable]] = await table.find_many( 

385 where={ 

386 "file_purpose": CONTAINER_OBJECT_PURPOSE, 

387 "created_by": {"in": owner_scopes}, 

388 } 

389 ) 

390 allowed_ids: Final = frozenset( 

391 row.model_object_id for row in rows if getattr(row, "model_object_id", None) is not None 

392 ) 

393 _ALLOWED_CONTAINER_IDS_CACHE.set_cache(cache_key, tuple(allowed_ids)) 

394 return allowed_ids 

395 

396 

397def _is_owned_container(item: object, allowed_container_ids: AbstractSet[str], custom_llm_provider: str) -> bool: 

398 container_id: Final = _get_response_id(item) 

399 if container_id is None: 

400 return False 

401 original_container_id, resolved_provider = decode_container_id_for_ownership(container_id, custom_llm_provider) 

402 return _container_model_object_id(original_container_id, resolved_provider) in allowed_container_ids 

403 

404 

405async def _collect_owned_containers( 

406 fetch_page: FetchContainerListPage, 

407 after: str | None, 

408 needed: int, 

409 allowed_container_ids: AbstractSet[str], 

410 custom_llm_provider: str, 

411 pages_left: int, 

412 collected: tuple[object, ...], 

413) -> tuple[object, tuple[object, ...]]: 

414 page: Final = await fetch_page(after, OWNED_CONTAINER_LIST_PAGE_SIZE) 

415 page_data: Final = _get_container_list_data(page) or () 

416 owned: Final = collected + tuple( 

417 item for item in page_data if _is_owned_container(item, allowed_container_ids, custom_llm_provider) 

418 ) 

419 upstream_last_id: Final = _get_response_id(page_data[-1]) if page_data else None 

420 if len(owned) >= needed or upstream_last_id is None or pages_left <= 1 or not _get_has_more(page): 

421 return page, owned 

422 return await _collect_owned_containers( 

423 fetch_page=fetch_page, 

424 after=upstream_last_id, 

425 needed=needed, 

426 allowed_container_ids=allowed_container_ids, 

427 custom_llm_provider=custom_llm_provider, 

428 pages_left=pages_left - 1, 

429 collected=owned, 

430 ) 

431 

432 

433async def list_owned_containers( 

434 fetch_page: FetchContainerListPage, 

435 after: str | None, 

436 limit: int | None, 

437 user_api_key_dict: UserAPIKeyAuth, 

438 custom_llm_provider: str, 

439) -> object: 

440 allowed_container_ids: Final = await _get_allowed_container_ids(user_api_key_dict) 

441 page_limit: Final = limit if limit is not None else DEFAULT_CONTAINER_LIST_LIMIT 

442 last_page, owned = await _collect_owned_containers( 

443 fetch_page=fetch_page, 

444 after=after, 

445 needed=page_limit + 1, 

446 allowed_container_ids=allowed_container_ids, 

447 custom_llm_provider=custom_llm_provider, 

448 pages_left=OWNED_CONTAINER_LIST_MAX_PAGES, 

449 collected=(), 

450 ) 

451 return _with_container_list_page( 

452 last_page, 

453 owned[:page_limit], 

454 has_more=len(owned) > page_limit or _get_has_more(last_page), 

455 )