80 lines
2.9 KiB
Python
80 lines
2.9 KiB
Python
from typing import Any
|
|||
|
|
|
||
|
|
from sqlalchemy import MetaData, Table, insert, select, update
|
||
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
|
|
||
|
|
TABLES = frozenset(
|
||
|
|
{
|
||
|
|
"config_release",
|
||
|
|
"platform_config_item",
|
||
|
|
"model_endpoint_config",
|
||
|
|
"model_routing_rule",
|
||
|
|
"model_routing_fallback",
|
||
|
|
"prompt_template_version",
|
||
|
|
"agent_intent_config",
|
||
|
|
"agent_reply_template",
|
||
|
|
"agent_negative_word",
|
||
|
|
"interaction_audit",
|
||
|
|
"api_request_receipt",
|
||
|
|
"svc_conversation_session",
|
||
|
|
"svc_handover_ticket",
|
||
|
|
"fin_knowledge_meta",
|
||
|
|
"profile_snapshots",
|
||
|
|
}
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
class PlatformRepository:
|
||
|
|
"""Core mappings reflected from frozen schema; resource names are a code whitelist."""
|
||
|
|
|
||
|
|
def __init__(self, session: AsyncSession) -> None:
|
||
|
|
self.session = session
|
||
|
|
self.tables: dict[str, Table] = {}
|
||
|
|
|
||
|
|
async def table(self, name: str) -> Table:
|
||
|
|
if name not in TABLES:
|
||
|
|
raise ValueError("unknown platform resource")
|
||
|
|
if name not in self.tables:
|
||
|
|
connection = await self.session.connection()
|
||
|
|
self.tables[name] = await connection.run_sync(
|
||
|
|
lambda sync: Table(name, MetaData(), autoload_with=sync)
|
||
|
|
)
|
||
|
|
return self.tables[name]
|
||
|
|
|
||
|
|
async def get(self, name: str, row_id: int, *, lock: bool = False) -> dict[str, Any] | None:
|
||
|
|
table = await self.table(name)
|
||
|
|
query = select(table).where(table.c.id == row_id)
|
||
|
|
if lock:
|
||
|
|
query = query.with_for_update()
|
||
|
|
row = (await self.session.execute(query)).mappings().first()
|
||
|
|
return dict(row) if row else None
|
||
|
|
|
||
|
|
async def rows(
|
||
|
|
self, name: str, filters: dict[str, Any], *, before: int | None = None, limit: int = 20
|
||
|
|
) -> list[dict[str, Any]]:
|
||
|
|
table = await self.table(name)
|
||
|
|
query = select(table)
|
||
|
|
for key, value in filters.items():
|
||
|
|
query = query.where(table.c[key] == value)
|
||
|
|
if before is not None:
|
||
|
|
query = query.where(table.c.id < before)
|
||
|
|
rows = (
|
||
|
|
await self.session.execute(query.order_by(table.c.id.desc()).limit(limit))
|
||
|
|
).mappings()
|
||
|
|
return [dict(row) for row in rows]
|
||
|
|
|
||
|
|
async def create(self, name: str, values: dict[str, Any]) -> dict[str, Any]:
|
||
|
|
table = await self.table(name)
|
||
|
|
result = await self.session.execute(insert(table).values(**values))
|
||
|
|
row_id = int(result.inserted_primary_key[0]) # type: ignore[attr-defined]
|
||
|
|
row = await self.get(name, row_id)
|
||
|
|
assert row is not None
|
||
|
|
return row
|
||
|
|
|
||
|
|
async def update(self, name: str, row_id: int, values: dict[str, Any]) -> dict[str, Any]:
|
||
|
|
table = await self.table(name)
|
||
|
|
await self.session.execute(update(table).where(table.c.id == row_id).values(**values))
|
||
|
|
row = await self.get(name, row_id)
|
||
|
|
assert row is not None
|
||
|
|
return row
|