@@ -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),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user