feat(cloud): manage LLM providers in database
Tests / Test passed: 664

This commit is contained in:
2026-07-14 00:31:47 +08:00
parent b613a315ff
commit a166ffd8a4
35 changed files with 2432 additions and 204 deletions
+355 -40
View File
@@ -20,6 +20,8 @@ from cloud.db_models import (
TaskAttemptRow,
AuthAuditRow,
HostGovernancePolicyRow,
LlmProviderProfileRow,
LlmProviderSettingsRow,
LoginThrottleRow,
UserRow,
UserSessionRow,
@@ -42,6 +44,19 @@ class SQLAlchemyCloudRepository:
self._sessions = sessionmaker(bind=engine, expire_on_commit=False)
if create_schema:
Base.metadata.create_all(engine)
self._ensure_llm_provider_settings()
def _ensure_llm_provider_settings(self) -> None:
with self._sessions.begin() as session:
if session.get(LlmProviderSettingsRow, "global") is None:
session.add(
LlmProviderSettingsRow(
id="global",
active_profile_id=None,
revision=0,
updated_at=_iso(utc_now()),
)
)
def enroll_host(
self,
@@ -503,7 +518,9 @@ class SQLAlchemyCloudRepository:
created_at=_iso(user.created_at),
updated_at=_iso(user.updated_at),
last_login_at=(
_iso(user.last_login_at) if user.last_login_at is not None else None
_iso(user.last_login_at)
if user.last_login_at is not None
else None
),
)
session.add(row)
@@ -576,7 +593,9 @@ class SQLAlchemyCloudRepository:
raise LastAdministratorConflictError(
"cannot remove the last enabled administrator"
)
security_changed = next_role != row.role or next_enabled != bool(row.enabled)
security_changed = next_role != row.role or next_enabled != bool(
row.enabled
)
if display_name is not None:
row.display_name = display_name
row.role = next_role
@@ -764,7 +783,9 @@ class SQLAlchemyCloudRepository:
session.flush()
return _login_throttle_from_row(row)
def clear_login_throttle(self, username_normalized: str, client_bucket: str) -> None:
def clear_login_throttle(
self, username_normalized: str, client_bucket: str
) -> None:
with self._sessions.begin() as session:
row = session.get(LoginThrottleRow, (username_normalized, client_bucket))
if row is not None:
@@ -843,7 +864,9 @@ class SQLAlchemyCloudRepository:
)
if row is None:
if expected_revision not in {None, 0}:
raise GovernancePolicyConflictError("submission policy revision changed")
raise GovernancePolicyConflictError(
"submission policy revision changed"
)
row = UserSubmissionPolicyRow(
user_id=user_id,
revision=1,
@@ -856,11 +879,10 @@ class SQLAlchemyCloudRepository:
)
session.add(row)
else:
if (
expected_revision is not None
and expected_revision != row.revision
):
raise GovernancePolicyConflictError("submission policy revision changed")
if expected_revision is not None and expected_revision != row.revision:
raise GovernancePolicyConflictError(
"submission policy revision changed"
)
row.revision += 1
row.submission_enabled = 1 if submission_enabled else 0
row.allowed_host_ids_json = _dump_optional_list(allowed_host_ids)
@@ -909,10 +931,7 @@ class SQLAlchemyCloudRepository:
)
session.add(row)
else:
if (
expected_revision is not None
and expected_revision != row.revision
):
if expected_revision is not None and expected_revision != row.revision:
raise GovernancePolicyConflictError("Host policy revision changed")
row.revision += 1
row.self_submission_enabled = 1 if self_submission_enabled else 0
@@ -922,6 +941,220 @@ class SQLAlchemyCloudRepository:
session.flush()
return _host_governance_policy_from_row(row)
# --------------------------------------------------------- LLM Providers
def get_llm_provider_settings(self) -> Any:
with self._sessions() as session:
row = session.get(LlmProviderSettingsRow, "global")
return _llm_provider_settings_from_row(row)
def list_llm_provider_profiles(self) -> list[Any]:
with self._sessions() as session:
rows = session.scalars(
select(LlmProviderProfileRow).order_by(
LlmProviderProfileRow.name_normalized
)
).all()
return [_llm_provider_profile_from_row(row) for row in rows]
def get_llm_provider_profile(self, profile_id: str) -> Any | None:
with self._sessions() as session:
row = session.get(LlmProviderProfileRow, profile_id)
return _llm_provider_profile_from_row(row) if row is not None else None
def create_llm_provider_profile(self, profile: Any) -> Any:
from cloud.llm_providers import LlmProviderConflictError
try:
with self._sessions.begin() as session:
if (
session.scalars(
select(LlmProviderProfileRow).where(
LlmProviderProfileRow.name_normalized
== profile.name_normalized
)
).first()
is not None
):
raise LlmProviderConflictError(
"Provider profile name already exists"
)
row = LlmProviderProfileRow(
id=profile.id,
name=profile.name,
name_normalized=profile.name_normalized,
provider_type=profile.provider_type,
model=profile.model,
base_url=profile.base_url,
timeout_seconds=profile.timeout_seconds,
api_key_ciphertext=profile.api_key_ciphertext,
key_last_rotated_at=_iso(profile.key_last_rotated_at),
enabled=1 if profile.enabled else 0,
revision=profile.revision,
created_at=_iso(profile.created_at),
updated_at=_iso(profile.updated_at),
)
session.add(row)
session.flush()
return _llm_provider_profile_from_row(row)
except IntegrityError as exc:
raise LlmProviderConflictError(
"Provider profile name already exists"
) from exc
def update_llm_provider_profile(
self,
profile_id: str,
*,
name: str,
name_normalized: str,
provider_type: str,
model: str,
base_url: str | None,
timeout_seconds: float,
api_key_ciphertext: str,
key_last_rotated_at: datetime,
enabled: bool,
expected_revision: int | None,
updated_at: datetime,
) -> Any:
from cloud.llm_providers import (
LlmProviderActiveConflictError,
LlmProviderConflictError,
)
try:
with self._sessions.begin() as session:
row = session.get(
LlmProviderProfileRow,
profile_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None:
raise KeyError(profile_id)
if expected_revision is not None and expected_revision != row.revision:
raise LlmProviderConflictError("Provider profile revision changed")
existing_name = session.scalars(
select(LlmProviderProfileRow).where(
LlmProviderProfileRow.name_normalized == name_normalized,
LlmProviderProfileRow.id != profile_id,
)
).first()
if existing_name is not None:
raise LlmProviderConflictError(
"Provider profile name already exists"
)
settings = session.get(
LlmProviderSettingsRow,
"global",
with_for_update=self.engine.dialect.name == "postgresql",
)
if (
not enabled
and settings is not None
and settings.active_profile_id == profile_id
):
raise LlmProviderActiveConflictError(
"activate another Provider profile before disabling this one"
)
row.name = name
row.name_normalized = name_normalized
row.provider_type = provider_type
row.model = model
row.base_url = base_url
row.timeout_seconds = timeout_seconds
row.api_key_ciphertext = api_key_ciphertext
row.key_last_rotated_at = _iso(key_last_rotated_at)
row.enabled = 1 if enabled else 0
row.revision += 1
row.updated_at = _iso(updated_at)
session.flush()
return _llm_provider_profile_from_row(row)
except IntegrityError as exc:
raise LlmProviderConflictError(
"Provider profile name already exists"
) from exc
def activate_llm_provider_profile(
self,
profile_id: str,
*,
expected_settings_revision: int | None,
updated_at: datetime,
) -> Any:
from cloud.llm_providers import LlmProviderConflictError
with self._sessions.begin() as session:
profile = session.get(
LlmProviderProfileRow,
profile_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if profile is None:
raise KeyError(profile_id)
if not profile.enabled:
raise LlmProviderConflictError(
"Provider profile must be enabled before activation"
)
settings = session.get(
LlmProviderSettingsRow,
"global",
with_for_update=self.engine.dialect.name == "postgresql",
)
if settings is None:
if expected_settings_revision not in {None, 0}:
raise LlmProviderConflictError("Provider settings revision changed")
settings = LlmProviderSettingsRow(
id="global",
active_profile_id=profile_id,
revision=1,
updated_at=_iso(updated_at),
)
session.add(settings)
else:
if (
expected_settings_revision is not None
and expected_settings_revision != settings.revision
):
raise LlmProviderConflictError("Provider settings revision changed")
settings.active_profile_id = profile_id
settings.revision += 1
settings.updated_at = _iso(updated_at)
session.flush()
return _llm_provider_settings_from_row(settings)
def delete_llm_provider_profile(
self,
profile_id: str,
*,
expected_revision: int | None,
) -> None:
from cloud.llm_providers import (
LlmProviderActiveConflictError,
LlmProviderConflictError,
)
with self._sessions.begin() as session:
row = session.get(
LlmProviderProfileRow,
profile_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None:
raise KeyError(profile_id)
if expected_revision is not None and expected_revision != row.revision:
raise LlmProviderConflictError("Provider profile revision changed")
settings = session.get(
LlmProviderSettingsRow,
"global",
with_for_update=self.engine.dialect.name == "postgresql",
)
if settings is not None and settings.active_profile_id == profile_id:
raise LlmProviderActiveConflictError(
"activate another Provider profile before deleting this one"
)
session.delete(row)
def count_active_tasks_for_host(self, host_id: str) -> int:
with self._sessions() as session:
count = session.scalar(
@@ -958,24 +1191,36 @@ class SQLAlchemyCloudRepository:
if policy is None or policy.daily_token_budget is None:
return None
used = session.scalar(
select(func.coalesce(func.sum(TokenUsageEventRow.total_tokens), 0)).where(
select(
func.coalesce(func.sum(TokenUsageEventRow.total_tokens), 0)
).where(
TokenUsageEventRow.host_id == host_id,
TokenUsageEventRow.usage_day == usage_day,
)
)
reserved = session.scalar(
select(func.coalesce(func.sum(TokenReservationRow.reserved_tokens), 0)).where(
select(
func.coalesce(func.sum(TokenReservationRow.reserved_tokens), 0)
).where(
TokenReservationRow.host_id == host_id,
TokenReservationRow.usage_day == usage_day,
TokenReservationRow.expires_at > _iso(created_at),
)
)
if int(used or 0) + int(reserved or 0) + reserved_tokens > policy.daily_token_budget:
if (
int(used or 0) + int(reserved or 0) + reserved_tokens
> policy.daily_token_budget
):
raise TokenBudgetExceededError("Host daily token budget is exhausted")
row = TokenReservationRow(
id=reservation_id, host_id=host_id, usage_day=usage_day,
reserved_tokens=reserved_tokens, task_id=task_id, attempt=attempt,
created_at=_iso(created_at), expires_at=_iso(expires_at),
id=reservation_id,
host_id=host_id,
usage_day=usage_day,
reserved_tokens=reserved_tokens,
task_id=task_id,
attempt=attempt,
created_at=_iso(created_at),
expires_at=_iso(expires_at),
)
session.add(row)
session.flush()
@@ -995,16 +1240,24 @@ class SQLAlchemyCloudRepository:
) -> Any | None:
with self._sessions.begin() as session:
row = session.get(
TokenReservationRow, reservation_id,
TokenReservationRow,
reservation_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None:
return None
event = TokenUsageEventRow(
id=event_id, host_id=row.host_id, usage_day=row.usage_day,
task_id=row.task_id, attempt=row.attempt, provider=provider, model=model,
input_tokens=input_tokens, output_tokens=output_tokens,
total_tokens=total_tokens, occurred_at=_iso(occurred_at),
id=event_id,
host_id=row.host_id,
usage_day=row.usage_day,
task_id=row.task_id,
attempt=row.attempt,
provider=provider,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
occurred_at=_iso(occurred_at),
)
session.add(event)
session.delete(row)
@@ -1024,39 +1277,55 @@ class SQLAlchemyCloudRepository:
return len(rows)
def get_host_token_usage_summary(
self, *, host_id: str, usage_day: str, now: datetime,
self,
*,
host_id: str,
usage_day: str,
now: datetime,
) -> Any:
from cloud.governance import TokenUsageSummary
with self._sessions() as session:
policy = session.get(HostGovernancePolicyRow, host_id)
used = session.scalar(
select(func.coalesce(func.sum(TokenUsageEventRow.total_tokens), 0)).where(
select(
func.coalesce(func.sum(TokenUsageEventRow.total_tokens), 0)
).where(
TokenUsageEventRow.host_id == host_id,
TokenUsageEventRow.usage_day == usage_day,
)
)
reserved = session.scalar(
select(func.coalesce(func.sum(TokenReservationRow.reserved_tokens), 0)).where(
select(
func.coalesce(func.sum(TokenReservationRow.reserved_tokens), 0)
).where(
TokenReservationRow.host_id == host_id,
TokenReservationRow.usage_day == usage_day,
TokenReservationRow.expires_at > _iso(now),
)
)
return TokenUsageSummary(
host_id=host_id, usage_day=usage_day,
host_id=host_id,
usage_day=usage_day,
daily_token_budget=(policy.daily_token_budget if policy else None),
used_tokens=int(used or 0), reserved_tokens=int(reserved or 0),
used_tokens=int(used or 0),
reserved_tokens=int(reserved or 0),
)
def list_host_token_usage_events(
self, *, host_id: str, limit: int, offset: int,
self,
*,
host_id: str,
limit: int,
offset: int,
) -> list[Any]:
with self._sessions() as session:
rows = session.scalars(
select(TokenUsageEventRow)
.where(TokenUsageEventRow.host_id == host_id)
.order_by(TokenUsageEventRow.occurred_at.desc(), TokenUsageEventRow.id.desc())
.order_by(
TokenUsageEventRow.occurred_at.desc(), TokenUsageEventRow.id.desc()
)
.limit(limit)
.offset(offset)
).all()
@@ -1423,7 +1692,9 @@ def _host_from_row(row: HostRow) -> Any:
address=row.address,
last_seen_at=_parse_dt(row.last_seen_at) or utc_now(),
planner_transport=(
row.planner_transport if row.planner_transport in {"direct", "cloud"} else "direct"
row.planner_transport
if row.planner_transport in {"direct", "cloud"}
else "direct"
),
)
@@ -1519,7 +1790,7 @@ def _load_optional_string_tuple(value: str | None) -> tuple[str, ...] | None:
return None
try:
parsed = json.loads(value)
except (TypeError, ValueError):
except TypeError, ValueError:
parsed = []
return tuple(item for item in parsed if isinstance(item, str))
@@ -1531,7 +1802,7 @@ def _load_optional_device_targets(
return None
try:
parsed = json.loads(value)
except (TypeError, ValueError):
except TypeError, ValueError:
parsed = []
return tuple(
(item[0], item[1])
@@ -1574,8 +1845,12 @@ def _token_reservation_from_row(row: TokenReservationRow) -> Any:
from cloud.governance import TokenReservation
return TokenReservation(
id=row.id, host_id=row.host_id, usage_day=row.usage_day,
reserved_tokens=row.reserved_tokens, task_id=row.task_id, attempt=row.attempt,
id=row.id,
host_id=row.host_id,
usage_day=row.usage_day,
reserved_tokens=row.reserved_tokens,
task_id=row.task_id,
attempt=row.attempt,
created_at=_parse_dt(row.created_at) or utc_now(),
expires_at=_parse_dt(row.expires_at) or utc_now(),
)
@@ -1585,10 +1860,17 @@ def _token_usage_event_from_row(row: TokenUsageEventRow) -> Any:
from cloud.governance import TokenUsageEvent
return TokenUsageEvent(
id=row.id, host_id=row.host_id, usage_day=row.usage_day,
task_id=row.task_id, attempt=row.attempt, provider=row.provider, model=row.model,
input_tokens=row.input_tokens, output_tokens=row.output_tokens,
total_tokens=row.total_tokens, occurred_at=_parse_dt(row.occurred_at) or utc_now(),
id=row.id,
host_id=row.host_id,
usage_day=row.usage_day,
task_id=row.task_id,
attempt=row.attempt,
provider=row.provider,
model=row.model,
input_tokens=row.input_tokens,
output_tokens=row.output_tokens,
total_tokens=row.total_tokens,
occurred_at=_parse_dt(row.occurred_at) or utc_now(),
)
@@ -1733,3 +2015,36 @@ def _plugin_from_row(row: PluginRow) -> tuple[Any, bool]:
),
bool(row.wired),
)
def _llm_provider_profile_from_row(row: LlmProviderProfileRow) -> Any:
from cloud.llm_providers import LlmProviderProfile
now = utc_now()
return LlmProviderProfile(
id=row.id,
name=row.name,
name_normalized=row.name_normalized,
provider_type=row.provider_type,
model=row.model,
base_url=row.base_url,
timeout_seconds=row.timeout_seconds,
api_key_ciphertext=row.api_key_ciphertext,
key_last_rotated_at=_parse_dt(row.key_last_rotated_at) or now,
enabled=bool(row.enabled),
revision=row.revision,
created_at=_parse_dt(row.created_at) or now,
updated_at=_parse_dt(row.updated_at) or now,
)
def _llm_provider_settings_from_row(row: LlmProviderSettingsRow | None) -> Any:
from cloud.llm_providers import LlmProviderSettings
if row is None:
return LlmProviderSettings(active_profile_id=None)
return LlmProviderSettings(
active_profile_id=row.active_profile_id,
revision=row.revision,
updated_at=_parse_dt(row.updated_at),
)