mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
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:
parent
f6b9518ddb
commit
b5ab0d15fe
8 changed files with 298 additions and 101 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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.")
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue