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>
This commit is contained in:
Devin AI 2026-08-10 13:32:51 +00:00
parent f6b9518ddb
commit b5ab0d15fe
8 changed files with 298 additions and 101 deletions

View file

@ -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

View file

@ -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.")

View file

@ -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:

View file

@ -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: ...

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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.