From b5ab0d15feee39c61948c7086cd32b8e0d37d8d5 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 10 Aug 2026 13:32:51 +0000 Subject: [PATCH] chore(typing): clear basedpyright Any errors in memory and tool management endpoints Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../tool_management_endpoints.py | 172 +++++++++++------- litellm/proxy/memory/memory_endpoints.py | 67 ++++--- .../object_permission_repository.py | 6 + litellm/repositories/prisma_protocols.py | 106 +++++++++++ litellm/repositories/table_repositories.py | 29 ++- litellm/repositories/team_repository.py | 6 + .../verification_token_repository.py | 6 + litellm/types/memory_management.py | 7 +- 8 files changed, 298 insertions(+), 101 deletions(-) diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 0fdedafb2bf..ae07acb5246 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -9,12 +9,14 @@ GET /v1/tool/{tool_name} - Get a single tool's details POST /v1/tool/policy - Update the input_policy / output_policy for a tool """ +import json import uuid +from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Annotated, Any, Final +from typing import TYPE_CHECKING, Annotated, Final from fastapi import APIRouter, Depends, HTTPException, Query -from pydantic import BaseModel, Field, TypeAdapter +from pydantic import BaseModel, Field, TypeAdapter, ValidationError if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient @@ -24,6 +26,7 @@ from litellm.constants import TOOL_SPEND_TOP_TOOLS from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.repositories.object_permission_repository import ObjectPermissionRepository +from litellm.repositories.prisma_protocols import DailyToolSpendRecord, SpendLogUsageRecord from litellm.repositories.table_repositories import ( DailyToolSpendRepository, SpendLogsRepository, @@ -201,9 +204,9 @@ async def get_tool_spend( end_str: Final = end_day.strftime("%Y-%m-%d") date_window: Final = {"date": {"gte": start_str, "lte": end_str}} - table: Final = DailyToolSpendRepository(prisma_client).table + rows: Final = DailyToolSpendRepository(prisma_client).rows top_tools: Final = _TOP_TOOL_ROWS.validate_python( - await table.group_by( + await rows.group_by( by=["tool_name"], sum={"spend": True, "total_tokens": True, "request_count": True}, where=date_window, @@ -222,8 +225,8 @@ async def get_tool_spend( for row in top_tools ] - daily_rows: Final = ( - await table.find_many( + daily_rows: Final[Sequence[DailyToolSpendRecord]] = ( + await rows.find_many( where={**date_window, "tool_name": {"in": [row.tool_name for row in top_tools]}}, order=[{"date": "asc"}, {"spend": "desc"}], ) @@ -270,54 +273,86 @@ async def get_tool_detail( raise HTTPException(status_code=500, detail=str(e)) -def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> str | None: +_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_JSON_ARRAY_ADAPTER: Final = TypeAdapter(list[object]) + + +def _as_json_object(value: object) -> Mapping[str, object] | None: + try: + return _JSON_OBJECT_ADAPTER.validate_python(value, strict=True) + except ValidationError: + return None + + +def _as_json_array(value: object) -> Sequence[object] | None: + try: + return _JSON_ARRAY_ADAPTER.validate_python(value, strict=True) + except ValidationError: + return None + + +def _decoded_json_or_raw(raw: str) -> object: + try: + decoded: Final[object] = json.loads(raw) + except Exception: + return raw + return decoded + + +def _messages_from_proxy_server_request(psr: Mapping[str, object]) -> object: + messages: Final = psr.get("messages") + if messages is not None: + return messages + body: Final = _as_json_object(psr.get("body")) + if body is not None: + return body.get("messages") + return None + + +def _input_snippet_for_tool_log(sl: SpendLogUsageRecord | None, max_len: int = 200) -> str | None: """Short snippet from messages or proxy_server_request for tool usage log row.""" if sl is None: return None - messages: Final = getattr(sl, "messages", None) + messages: Final = sl.messages if messages is not None: - s = _snippet_str(messages, max_len) - if s: - return s - psr = getattr(sl, "proxy_server_request", None) + message_snippet: Final = _snippet_str(messages, max_len) + if message_snippet: + return message_snippet + psr: Final = sl.proxy_server_request if not psr: return None - if isinstance(psr, str): - import json - - try: - psr = json.loads(psr) - except Exception: - return _snippet_str(psr, max_len) - if isinstance(psr, dict): - msgs = psr.get("messages") - if msgs is None and isinstance(psr.get("body"), dict): - msgs = psr["body"].get("messages") - s = _snippet_str(msgs, max_len) - if s: - return s - return _snippet_str(psr, max_len) + decoded: Final = _decoded_json_or_raw(psr) if isinstance(psr, str) else psr + request_body: Final = _as_json_object(decoded) + if request_body is not None: + request_snippet: Final = _snippet_str(_messages_from_proxy_server_request(request_body), max_len) + if request_snippet: + return request_snippet + return _snippet_str(decoded, max_len) -def _snippet_str(text: Any, max_len: int = 200) -> str | None: +def _content_part_str(item: object) -> str: + entry: Final = _as_json_object(item) + if entry is not None and "content" in entry: + content: Final = entry["content"] + return content if isinstance(content, str) else str(content) + return str(item) + + +def _truncated_snippet(text: str, max_len: int) -> str | None: + if not text or text == "{}": + return None + return (text[:max_len] + "...") if len(text) > max_len else text + + +def _snippet_str(text: object, max_len: int = 200) -> str | None: if text is None: return None if isinstance(text, str): - s = text - elif isinstance(text, list): - parts: Final = [] - for item in text: - if isinstance(item, dict) and "content" in item: - c = item["content"] - parts.append(c if isinstance(c, str) else str(c)) - else: - parts.append(str(item)) - s = " ".join(parts) - else: - s = str(text) - if not s or s == "{}": - return None - return (s[:max_len] + "...") if len(s) > max_len else s + return _truncated_snippet(text, max_len) + items: Final = _as_json_array(text) + if items is not None: + return _truncated_snippet(" ".join(_content_part_str(item) for item in items), max_len) + return _truncated_snippet(str(text), max_len) @router.get( @@ -344,7 +379,7 @@ async def get_tool_usage_logs( raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) try: - where: Final[dict] = {"tool_name": tool_name} + where: Final[dict[str, object]] = {"tool_name": tool_name} if start_date or end_date: start_time_filter: datetime | None = None end_time_filter: datetime | None = None @@ -363,14 +398,15 @@ async def get_tool_usage_logs( except ValueError: pass if start_time_filter is not None or end_time_filter is not None: - where["start_time"] = {} + start_time_window: Final[dict[str, datetime]] = {} + where["start_time"] = start_time_window if start_time_filter is not None: - where["start_time"]["gte"] = start_time_filter + start_time_window["gte"] = start_time_filter if end_time_filter is not None: - where["start_time"]["lte"] = end_time_filter + start_time_window["lte"] = end_time_filter - total: Final = await SpendLogToolIndexRepository(prisma_client).table.count(where=where) - index_rows: Final = await SpendLogToolIndexRepository(prisma_client).table.find_many( + total: Final = await SpendLogToolIndexRepository(prisma_client).rows.count(where=where) + index_rows: Final = await SpendLogToolIndexRepository(prisma_client).rows.find_many( where=where, order={"start_time": "desc"}, skip=(page - 1) * page_size, @@ -380,7 +416,7 @@ async def get_tool_usage_logs( if not request_ids: return ToolUsageLogsResponse(logs=[], total=total, page=page, page_size=page_size) - spend_logs = await SpendLogsRepository(prisma_client).table.find_many(where={"request_id": {"in": request_ids}}) + spend_logs: Final = await SpendLogsRepository(prisma_client).find_usage_by_request_ids(request_ids) log_by_id: Final = {s.request_id: s for s in spend_logs} logs_out: Final[list[ToolUsageLogEntry]] = [] @@ -449,24 +485,24 @@ async def _resolve_key_hash_to_object_permission_id( hashed: Final = key_hash if "sk-" not in (key_hash or "") else hash_token(key_hash) if not hashed: return None - row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed}) + token_rows: Final = VerificationTokenRepository(prisma_client).object_permission_rows + row: Final = await token_rows.find_unique(where={"token": hashed}) if row is None: return None - op_id: Final = getattr(row, "object_permission_id", None) + op_id: Final = row.object_permission_id if op_id: return op_id new_id: Final = str(uuid.uuid4()) - await ObjectPermissionRepository(prisma_client).table.create( - data={"object_permission_id": new_id, "blocked_tools": []} - ) - updated_count: Final = await VerificationTokenRepository(prisma_client).table.update_many( + permission_rows: Final = ObjectPermissionRepository(prisma_client).rows + await permission_rows.create(data={"object_permission_id": new_id, "blocked_tools": []}) + updated_count: Final = await token_rows.update_many( where={"token": hashed, "object_permission_id": None}, data={"object_permission_id": new_id}, ) if updated_count == 0: - await ObjectPermissionRepository(prisma_client).table.delete(where={"object_permission_id": new_id}) - row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed}) - return getattr(row, "object_permission_id", None) if row else None + await permission_rows.delete(where={"object_permission_id": new_id}) + row_after_race: Final = await token_rows.find_unique(where={"token": hashed}) + return row_after_race.object_permission_id if row_after_race else None return new_id @@ -478,24 +514,24 @@ async def _resolve_team_id_to_object_permission_id( if not team_id or not team_id.strip(): return None team_id_clean: Final = team_id.strip() - row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id_clean}) + team_rows: Final = TeamRepository(prisma_client).object_permission_rows + row: Final = await team_rows.find_unique(where={"team_id": team_id_clean}) if row is None: return None - op_id: Final = getattr(row, "object_permission_id", None) + op_id: Final = row.object_permission_id if op_id: return op_id new_id: Final = str(uuid.uuid4()) - await ObjectPermissionRepository(prisma_client).table.create( - data={"object_permission_id": new_id, "blocked_tools": []} - ) - updated_count: Final = await TeamRepository(prisma_client).table.update_many( + permission_rows: Final = ObjectPermissionRepository(prisma_client).rows + await permission_rows.create(data={"object_permission_id": new_id, "blocked_tools": []}) + updated_count: Final = await team_rows.update_many( where={"team_id": team_id_clean, "object_permission_id": None}, data={"object_permission_id": new_id}, ) if updated_count == 0: - await ObjectPermissionRepository(prisma_client).table.delete(where={"object_permission_id": new_id}) - row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id_clean}) - return getattr(row, "object_permission_id", None) if row else None + await permission_rows.delete(where={"object_permission_id": new_id}) + row_after_race: Final = await team_rows.find_unique(where={"team_id": team_id_clean}) + return row_after_race.object_permission_id if row_after_race else None return new_id diff --git a/litellm/proxy/memory/memory_endpoints.py b/litellm/proxy/memory/memory_endpoints.py index 987823d987f..31520a9a516 100644 --- a/litellm/proxy/memory/memory_endpoints.py +++ b/litellm/proxy/memory/memory_endpoints.py @@ -18,7 +18,8 @@ Scoping: """ import json -from typing import Any, Final +from collections.abc import Mapping +from typing import TYPE_CHECKING, Final from fastapi import APIRouter, Depends, HTTPException, Query @@ -30,6 +31,7 @@ from litellm.proxy._types import ( user_api_key_has_admin_view, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.repositories.prisma_protocols import MemoryRecord from litellm.repositories.table_repositories import MemoryRepository from litellm.repositories.team_repository import TeamRepository from litellm.types.memory_management import ( @@ -40,14 +42,17 @@ from litellm.types.memory_management import ( MemoryUpdateRequest, ) +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + router: Final = APIRouter() -def _serialize_metadata_for_prisma(metadata: Any) -> str: +def _serialize_metadata_for_prisma(metadata: object) -> str: """ Encode a `metadata` payload for the `Json?` column. - `metadata` is typed `Optional[Any]`, so callers may send dicts, lists, + `metadata` is typed `object | None`, so callers may send dicts, lists, or JSON scalars (including plain Python strings like `"hello"`). prisma-client-python rejects raw Python values on `Json?` columns (`MissingRequiredValueError` / `DataError`), and Postgres `jsonb` @@ -62,14 +67,14 @@ def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN -def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> dict | None: +def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, object] | None: """ Prisma `where` fragment restricting rows to those the caller can see. Returns None for admins (no restriction). """ if user_api_key_has_admin_view(user_api_key_dict): return None - ors: Final[list[dict]] = [] + ors: Final[list[dict[str, object]]] = [] if user_api_key_dict.user_id: ors.append({"user_id": user_api_key_dict.user_id}) if user_api_key_dict.team_id: @@ -80,7 +85,7 @@ def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> dict | None: return {"OR": ors} -def _row_to_model(row: Any) -> LiteLLM_MemoryRow: +def _row_to_model(row: MemoryRecord) -> LiteLLM_MemoryRow: return LiteLLM_MemoryRow( memory_id=row.memory_id, key=row.key, @@ -95,7 +100,7 @@ def _row_to_model(row: Any) -> LiteLLM_MemoryRow: ) -def _require_prisma(): +def _require_prisma() -> "PrismaClient": from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -113,7 +118,9 @@ def _internal_error(log_message: str, exc: Exception, default_detail: str) -> HT return HTTPException(status_code=500, detail=default_detail) -async def _assert_write_access(prisma_client: Any, row: Any, user_api_key_dict: UserAPIKeyAuth) -> None: +async def _assert_write_access( + prisma_client: "PrismaClient", row: MemoryRecord, user_api_key_dict: UserAPIKeyAuth +) -> None: """ Enforce ownership for mutations (PUT/DELETE). @@ -135,8 +142,8 @@ async def _assert_write_access(prisma_client: Any, row: Any, user_api_key_dict: """ if _is_admin(user_api_key_dict): return - row_user_id: Final = getattr(row, "user_id", None) - row_team_id: Final = getattr(row, "team_id", None) + row_user_id: Final = row.user_id + row_team_id: Final = row.team_id # Personal ownership. if row_user_id and row_user_id == user_api_key_dict.user_id: @@ -153,7 +160,7 @@ async def _assert_write_access(prisma_client: Any, row: Any, user_api_key_dict: ) -async def _is_team_admin_for(prisma_client: Any, user_api_key_dict: UserAPIKeyAuth, team_id: str) -> bool: +async def _is_team_admin_for(prisma_client: "PrismaClient", user_api_key_dict: UserAPIKeyAuth, team_id: str) -> bool: """ True if the caller is a team admin of `team_id`, or an org admin for the team's organization. Mirrors the auth pattern used by team-management @@ -269,7 +276,7 @@ async def create_memory( # `metadata` is a `Json?` column — prisma-client-python rejects raw # Python values, so JSON-encode any non-null payload and omit the field # entirely when None so the column defaults to SQL NULL. - create_data: Final[dict] = { + create_data: Final[dict[str, object]] = { "key": body.key, "value": body.value, "user_id": user_id, @@ -281,7 +288,7 @@ async def create_memory( create_data["metadata"] = _serialize_metadata_for_prisma(body.metadata) try: - row: Final = await MemoryRepository(prisma_client).table.create(data=create_data) + row: Final = await MemoryRepository(prisma_client).rows.create(data=create_data) except Exception as e: # Key is globally unique. Any duplicate → 409. if _is_unique_violation(e): @@ -325,14 +332,14 @@ async def list_memory( # top-level "AND" — safer than `dict.update` since future visibility # filters could grow an "OR" key that would clobber this one if merged # by key. - key_filter: Final[dict] = {} + key_filter: Final[dict[str, object]] = {} if key_prefix is not None: key_filter["key"] = {"startsWith": key_prefix} elif key is not None: key_filter["key"] = key vis: Final = _visibility_filter(user_api_key_dict) - where: dict + where: Mapping[str, object] if vis is None: where = key_filter elif not key_filter: @@ -341,8 +348,8 @@ async def list_memory( where = {"AND": [key_filter, vis]} try: - total: Final = await MemoryRepository(prisma_client).table.count(where=where) - rows: Final = await MemoryRepository(prisma_client).table.find_many( + total: Final = await MemoryRepository(prisma_client).rows.count(where=where) + rows: Final = await MemoryRepository(prisma_client).rows.find_many( where=where, order={"updated_at": "desc"}, skip=(page - 1) * page_size, @@ -354,12 +361,16 @@ async def list_memory( return MemoryListResponse(memories=[_row_to_model(r) for r in rows], total=total) -async def _find_memory_for_caller(prisma_client: Any, key: str, user_api_key_dict: UserAPIKeyAuth) -> Any: +async def _find_memory_for_caller( + prisma_client: "PrismaClient", key: str, user_api_key_dict: UserAPIKeyAuth +) -> MemoryRecord: """Look up a memory row by key, scoped to the caller's visibility.""" - key_filter: Final[dict] = {"key": key} + key_filter: Final[Mapping[str, object]] = {"key": key} vis: Final = _visibility_filter(user_api_key_dict) - where: Final[dict] = key_filter if vis is None else {"AND": [key_filter, vis]} - rows = await MemoryRepository(prisma_client).table.find_many(where=where, take=1, order={"updated_at": "desc"}) + where: Final[Mapping[str, object]] = key_filter if vis is None else {"AND": [key_filter, vis]} + rows: Final = await MemoryRepository(prisma_client).rows.find_many( + where=where, take=1, order={"updated_at": "desc"} + ) if not rows: raise HTTPException(status_code=404, detail=f"Memory with key '{key}' not found") return rows[0] @@ -415,7 +426,7 @@ async def upsert_memory( fields_sent: Final = body.model_fields_set metadata_in_payload: Final = "metadata" in fields_sent - data: Final[dict] = {} + data: Final[dict[str, object]] = {} if body.value is not None: data["value"] = body.value if metadata_in_payload: @@ -427,7 +438,7 @@ async def upsert_memory( ) data["updated_by"] = user_api_key_dict.user_id - async def _find_existing() -> Any: + async def _find_existing() -> MemoryRecord | None: """Return the caller-visible row for `key`, or None.""" try: return await _find_memory_for_caller(prisma_client, key, user_api_key_dict) @@ -444,7 +455,7 @@ async def upsert_memory( # their team) — otherwise a teammate could overwrite a personal # entry through the OR-based visibility filter. await _assert_write_access(prisma_client, existing, user_api_key_dict) - row = await MemoryRepository(prisma_client).table.update( + row = await MemoryRepository(prisma_client).rows.update( where={"memory_id": existing.memory_id}, data=data, ) @@ -459,7 +470,7 @@ async def upsert_memory( # Omit `metadata` when None so the column defaults to SQL NULL; # otherwise JSON-encode for Prisma — same pattern as # `create_memory` above. - create_data: Final[dict] = { + create_data: Final[dict[str, object]] = { "key": key, "value": body.value, "user_id": user_id, @@ -470,7 +481,7 @@ async def upsert_memory( if body.metadata is not None: create_data["metadata"] = _serialize_metadata_for_prisma(body.metadata) try: - row = await MemoryRepository(prisma_client).table.create(data=create_data) + row = await MemoryRepository(prisma_client).rows.create(data=create_data) except Exception as e: # Race: a concurrent PUT/POST created the row after our check. # Re-read and fall back to an update so the PUT stays idempotent @@ -487,7 +498,7 @@ async def upsert_memory( ) # Same write-authorization check as the non-race path. await _assert_write_access(prisma_client, existing_after_race, user_api_key_dict) - row = await MemoryRepository(prisma_client).table.update( + row = await MemoryRepository(prisma_client).rows.update( where={"memory_id": existing_after_race.memory_id}, data=data, ) @@ -515,7 +526,7 @@ async def delete_memory( # Visibility != write authority — see the upsert handler for the rationale. await _assert_write_access(prisma_client, row, user_api_key_dict) try: - await MemoryRepository(prisma_client).table.delete(where={"memory_id": row.memory_id}) + await MemoryRepository(prisma_client).rows.delete(where={"memory_id": row.memory_id}) except Exception as e: raise _internal_error("Error deleting memory: %s", e, "Internal error deleting memory entry.") diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index 54a311c4a77..2b6aa4a9cb3 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -6,6 +6,7 @@ from typing import Any, Final from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.prisma_protocols import ObjectPermissionWriteTable class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): @@ -19,6 +20,11 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): def model_class(self) -> type[LiteLLM_ObjectPermissionTable]: return LiteLLM_ObjectPermissionTable + @property + def rows(self) -> ObjectPermissionWriteTable: + rows: Final[ObjectPermissionWriteTable] = self.table + return rows + async def find_by_id( self, object_permission_id: str, id_field: str = "object_permission_id" ) -> LiteLLM_ObjectPermissionTable | None: diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 6aff196ff10..3ee7db73a0b 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -7,6 +7,7 @@ private ones per file. """ from collections.abc import Mapping, Sequence +from datetime import datetime from typing import Protocol, TypeVar RowT_co = TypeVar("RowT_co", covariant=True) @@ -26,6 +27,111 @@ class SpendLinkedTable(Protocol[RowT_co]): async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... +class MemoryRecord(Protocol): + memory_id: str + key: str + value: str + metadata: object | None + user_id: str | None + team_id: str | None + created_at: datetime | None + created_by: str | None + updated_at: datetime | None + updated_by: str | None + + +class MemoryTable(Protocol): + async def create(self, *, data: Mapping[str, object]) -> MemoryRecord: ... + + async def count(self, *, where: Mapping[str, object]) -> int: ... + + async def find_many( + self, + *, + where: Mapping[str, object], + order: Mapping[str, str], + skip: int = 0, + take: int | None = None, + ) -> Sequence[MemoryRecord]: ... + + async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> MemoryRecord: ... + + async def delete(self, *, where: Mapping[str, object]) -> MemoryRecord | None: ... + + +class ToolIndexRecord(Protocol): + request_id: str + + +class ToolIndexTable(Protocol): + async def count(self, *, where: Mapping[str, object]) -> int: ... + + async def find_many( + self, + *, + where: Mapping[str, object], + order: Mapping[str, str], + skip: int = 0, + take: int | None = None, + ) -> Sequence[ToolIndexRecord]: ... + + +class SpendLogUsageRecord(Protocol): + request_id: str + startTime: datetime + model: str | None + spend: float | None + total_tokens: int | None + messages: object | None + proxy_server_request: object | None + + +class SpendLogUsageTable(Protocol): + async def find_many(self, *, where: Mapping[str, object]) -> Sequence[SpendLogUsageRecord]: ... + + +class DailyToolSpendRecord(Protocol): + date: str + tool_name: str + spend: float + request_count: int + + +class DailyToolSpendTable(Protocol): + async def group_by( + self, + *, + by: Sequence[str], + sum: Mapping[str, bool], + where: Mapping[str, object], + order: Mapping[str, object], + take: int, + ) -> Sequence[Mapping[str, object]] | None: ... + + async def find_many( + self, + *, + where: Mapping[str, object], + order: Sequence[Mapping[str, str]], + ) -> Sequence[DailyToolSpendRecord]: ... + + +class ObjectPermissionOwnerRecord(Protocol): + object_permission_id: str | None + + +class ObjectPermissionOwnerTable(Protocol): + async def find_unique(self, *, where: Mapping[str, object]) -> ObjectPermissionOwnerRecord | None: ... + + async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... + + +class ObjectPermissionWriteTable(Protocol): + async def create(self, *, data: Mapping[str, object]) -> object: ... + + async def delete(self, *, where: Mapping[str, object]) -> object: ... + + class BatchTable(Protocol): def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index be19f290ba6..2850cb5975c 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -7,9 +7,17 @@ These are thin wrappers for tables that do not (yet) need domain-specific query methods; richer repositories live in their own modules. """ -from typing import Any +from collections.abc import Sequence +from typing import Any, Final from litellm.proxy.common_utils.config_sync_pubsub import wrap_table_actions_for_config_sync +from litellm.repositories.prisma_protocols import ( + DailyToolSpendTable, + MemoryTable, + SpendLogUsageRecord, + SpendLogUsageTable, + ToolIndexTable, +) class PrismaTableRepository: @@ -65,6 +73,10 @@ class OrganizationMembershipRepository(PrismaTableRepository): class SpendLogsRepository(PrismaTableRepository): table_name = "litellm_spendlogs" + async def find_usage_by_request_ids(self, request_ids: Sequence[str]) -> Sequence[SpendLogUsageRecord]: + rows: Final[SpendLogUsageTable] = self.table + return await rows.find_many(where={"request_id": {"in": request_ids}}) + class ClaudeCodePluginRepository(PrismaTableRepository): table_name = "litellm_claudecodeplugintable" @@ -113,6 +125,11 @@ class ManagedFileRepository(PrismaTableRepository): class MemoryRepository(PrismaTableRepository): table_name = "litellm_memorytable" + @property + def rows(self) -> MemoryTable: + rows: Final[MemoryTable] = self.table + return rows + class SearchToolsRepository(PrismaTableRepository): table_name = "litellm_searchtoolstable" @@ -189,10 +206,20 @@ class DailyTagSpendRepository(PrismaTableRepository): class SpendLogToolIndexRepository(PrismaTableRepository): table_name = "litellm_spendlogtoolindex" + @property + def rows(self) -> ToolIndexTable: + rows: Final[ToolIndexTable] = self.table + return rows + class DailyToolSpendRepository(PrismaTableRepository): table_name = "litellm_dailytoolspend" + @property + def rows(self) -> DailyToolSpendTable: + rows: Final[DailyToolSpendTable] = self.table + return rows + class SpendLogGuardrailIndexRepository(PrismaTableRepository): table_name = "litellm_spendlogguardrailindex" diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index a4908af561a..01d08bfe0e4 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -15,6 +15,7 @@ from litellm.repositories.base_repository import ( DbRecord, record_to_dict, ) +from litellm.repositories.prisma_protocols import ObjectPermissionOwnerTable if TYPE_CHECKING: from prisma import Prisma @@ -41,6 +42,11 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): def deleted_table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper return self.prisma_client.db.litellm_deletedteamtable + @property + def object_permission_rows(self) -> ObjectPermissionOwnerTable: + rows: Final[ObjectPermissionOwnerTable] = self.table + return rows + @property def model_class(self) -> type[LiteLLM_TeamTable]: return LiteLLM_TeamTable diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index 3790ad25914..932b0162d7d 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -15,6 +15,7 @@ from litellm.repositories.base_repository import ( DbRecord, record_to_dict, ) +from litellm.repositories.prisma_protocols import ObjectPermissionOwnerTable if TYPE_CHECKING: from prisma.models import ( @@ -52,6 +53,11 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): def deleted_table(self) -> Any: return self.prisma_client.db.litellm_deletedverificationtoken + @property + def object_permission_rows(self) -> ObjectPermissionOwnerTable: + rows: Final[ObjectPermissionOwnerTable] = self.table + return rows + @property def model_class(self) -> type[LiteLLM_VerificationToken]: return LiteLLM_VerificationToken diff --git a/litellm/types/memory_management.py b/litellm/types/memory_management.py index 04a2a0c1905..153de0c6cb9 100644 --- a/litellm/types/memory_management.py +++ b/litellm/types/memory_management.py @@ -3,7 +3,6 @@ Pydantic models for Memory management endpoints. """ from datetime import datetime -from typing import Any from pydantic import BaseModel, Field @@ -12,7 +11,7 @@ class LiteLLM_MemoryRow(BaseModel): memory_id: str key: str value: str - metadata: Any | None = None + metadata: object | None = None user_id: str | None = None team_id: str | None = None created_at: datetime | None = None @@ -24,7 +23,7 @@ class LiteLLM_MemoryRow(BaseModel): class MemoryCreateRequest(BaseModel): key: str = Field(..., description="Memory key (acts as the namespace in the URL).") value: str = Field(..., description="Memory content. Typically markdown/text for LLM context.") - metadata: Any | None = Field( + metadata: object | None = Field( default=None, description="Optional JSON metadata (tags, structured fields).", ) @@ -40,7 +39,7 @@ class MemoryCreateRequest(BaseModel): class MemoryUpdateRequest(BaseModel): value: str | None = None - metadata: Any | None = None + metadata: object | None = None # Only honored on create (when the row doesn't yet exist) and only for # PROXY_ADMIN callers — mirrors MemoryCreateRequest so admins can bootstrap # rows scoped to another user/team via PUT, not just POST.