Fix auth and config validation regressions

This commit is contained in:
Abhishek Kumar 2026-07-09 18:18:44 +05:30
parent e4b53f78e9
commit 8ac6f1aa06
9 changed files with 346 additions and 49 deletions

View file

@ -1,3 +1,4 @@
from datetime import UTC, datetime
from typing import Any, Dict, List, Optional
from sqlalchemy.future import select
@ -66,6 +67,30 @@ class OrganizationConfigurationClient(BaseDBClient):
await session.refresh(config)
return config
async def touch_configuration(
self, organization_id: int, key: str
) -> Optional[OrganizationConfigurationModel]:
"""Update the timestamp for an existing organization configuration."""
async with self.async_session() as session:
result = await session.execute(
select(OrganizationConfigurationModel).where(
OrganizationConfigurationModel.organization_id == organization_id,
OrganizationConfigurationModel.key == key,
)
)
config = result.scalars().first()
if not config:
return None
config.updated_at = datetime.now(UTC)
try:
await session.commit()
except Exception as e:
await session.rollback()
raise e
await session.refresh(config)
return config
async def delete_configuration(self, organization_id: int, key: str) -> bool:
"""Delete a configuration for an organization."""
async with self.async_session() as session:

View file

@ -18,6 +18,7 @@ from api.services.auth.depends import get_user
from api.services.configuration.ai_model_configuration import (
convert_legacy_ai_model_configuration_to_v2,
get_resolved_ai_model_configuration,
update_organization_ai_model_configuration_last_validated_at,
upsert_organization_ai_model_configuration_v2,
)
from api.services.configuration.check_validity import (
@ -107,6 +108,24 @@ class UserConfigurationRequestResponseSchema(BaseModel):
organization_pricing: dict[str, Union[float, str, bool]] | None = None
def _is_validation_cache_stale(
last_validated_at: datetime | None,
validity_ttl_seconds: int,
) -> bool:
if last_validated_at is None:
return True
has_timezone = (
last_validated_at.tzinfo is not None
and last_validated_at.utcoffset() is not None
)
if has_timezone:
now = datetime.now(last_validated_at.tzinfo)
else:
now = datetime.now()
return last_validated_at < now - timedelta(seconds=validity_ttl_seconds)
@router.get("/configurations/user")
async def get_user_configurations(
user: UserModel = Depends(get_user),
@ -255,10 +274,9 @@ async def validate_user_configurations(
)
configurations = resolved_config.effective
if (
not configurations.last_validated_at
or configurations.last_validated_at
< datetime.now() - timedelta(seconds=validity_ttl_seconds)
if _is_validation_cache_stale(
configurations.last_validated_at,
validity_ttl_seconds,
):
validator = UserConfigurationValidator()
try:
@ -267,6 +285,13 @@ async def validate_user_configurations(
organization_id=user.selected_organization_id,
created_by=user.provider_id,
)
if (
resolved_config.source == "organization_v2"
and user.selected_organization_id is not None
):
await update_organization_ai_model_configuration_last_validated_at(
user.selected_organization_id
)
return status
except ValueError as e:
raise HTTPException(status_code=422, detail=e.args[0])

View file

@ -11,7 +11,11 @@ from sqlalchemy.orm import selectinload
from api.constants import MPS_API_URL
from api.db import db_client
from api.db.models import WorkflowDefinitionModel, WorkflowModel
from api.db.models import (
OrganizationConfigurationModel,
WorkflowDefinitionModel,
WorkflowModel,
)
from api.enums import OrganizationConfigurationKey
from api.schemas.ai_model_configuration import (
DOGRAH_DEFAULT_LANGUAGE,
@ -58,12 +62,19 @@ async def get_resolved_ai_model_configuration(
organization_id: int | None,
) -> ResolvedAIModelConfiguration:
"""Resolve the effective model configuration for an organization."""
organization_configuration = await get_organization_ai_model_configuration_v2(
organization_id
organization_configuration_row = (
await _get_organization_ai_model_configuration_v2_row(organization_id)
)
organization_configuration = _parse_organization_ai_model_configuration_v2(
organization_configuration_row,
organization_id,
)
if organization_configuration is not None:
effective = compile_ai_model_configuration_v2(organization_configuration)
if organization_configuration_row is not None:
effective.last_validated_at = organization_configuration_row.updated_at
return ResolvedAIModelConfiguration(
effective=compile_ai_model_configuration_v2(organization_configuration),
effective=effective,
source="organization_v2",
organization_configuration=organization_configuration,
)
@ -100,12 +111,34 @@ async def get_effective_ai_model_configuration_for_workflow(
async def get_organization_ai_model_configuration_v2(
organization_id: int | None,
) -> OrganizationAIModelConfigurationV2 | None:
if organization_id is None:
return None
row = await db_client.get_configuration(
row = await _get_organization_ai_model_configuration_v2_row(organization_id)
return _parse_organization_ai_model_configuration_v2(row, organization_id)
async def update_organization_ai_model_configuration_last_validated_at(
organization_id: int,
) -> None:
await db_client.touch_configuration(
organization_id,
OrganizationConfigurationKey.MODEL_CONFIGURATION_V2.value,
)
async def _get_organization_ai_model_configuration_v2_row(
organization_id: int | None,
) -> OrganizationConfigurationModel | None:
if organization_id is None:
return None
return await db_client.get_configuration(
organization_id,
OrganizationConfigurationKey.MODEL_CONFIGURATION_V2.value,
)
def _parse_organization_ai_model_configuration_v2(
row: OrganizationConfigurationModel | None,
organization_id: int | None,
) -> OrganizationAIModelConfigurationV2 | None:
if row is None or not row.value:
return None
try:

View file

@ -15,6 +15,7 @@ from api.services.configuration.ai_model_configuration import (
WORKFLOW_MODEL_CONFIGURATION_V2_OVERRIDE_KEY,
check_for_masked_keys_in_ai_model_configuration_v2,
convert_legacy_ai_model_configuration_to_v2,
get_resolved_ai_model_configuration,
mask_ai_model_configuration_v2,
merge_ai_model_configuration_v2_secrets,
migrate_workflow_configuration_model_override_to_v2,
@ -148,6 +149,35 @@ async def test_byok_realtime_validator_does_not_require_stt_or_tts():
}
@pytest.mark.asyncio
async def test_resolved_org_v2_uses_configuration_updated_at_as_validation_cache(
monkeypatch,
):
from datetime import UTC, datetime
from api.services.configuration import ai_model_configuration
last_validated_at = datetime.now(UTC)
config = OrganizationAIModelConfigurationV2(
mode="dograh",
dograh=DograhManagedAIModelConfiguration(api_key="mps-secret"),
)
row = SimpleNamespace(
value=config.model_dump(mode="json", exclude_none=True),
updated_at=last_validated_at,
)
monkeypatch.setattr(
ai_model_configuration.db_client,
"get_configuration",
AsyncMock(return_value=row),
)
resolved = await get_resolved_ai_model_configuration(organization_id=42)
assert resolved.source == "organization_v2"
assert resolved.effective.last_validated_at == last_validated_at
@pytest.mark.asyncio
async def test_pipeline_validator_requires_stt_and_tts_when_not_realtime():
effective = EffectiveAIModelConfiguration(

View file

@ -0,0 +1,99 @@
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from api.routes import user as user_routes
from api.schemas.ai_model_configuration import EffectiveAIModelConfiguration
from api.services.configuration.ai_model_configuration import (
ResolvedAIModelConfiguration,
)
@pytest.mark.asyncio
async def test_validate_user_configurations_marks_stale_org_v2_config_validated(
monkeypatch,
):
stale_config = EffectiveAIModelConfiguration(
last_validated_at=datetime.now(UTC) - timedelta(seconds=120)
)
resolved = ResolvedAIModelConfiguration(
effective=stale_config,
source="organization_v2",
)
validate = AsyncMock(return_value={"status": [{"model": "all", "message": "ok"}]})
touch_validation_cache = AsyncMock()
class FakeValidator:
def __init__(self):
self.validate = validate
monkeypatch.setattr(
user_routes,
"get_resolved_ai_model_configuration",
AsyncMock(return_value=resolved),
)
monkeypatch.setattr(user_routes, "UserConfigurationValidator", FakeValidator)
monkeypatch.setattr(
user_routes,
"update_organization_ai_model_configuration_last_validated_at",
touch_validation_cache,
)
response = await user_routes.validate_user_configurations(
validity_ttl_seconds=60,
user=SimpleNamespace(
provider_id="provider-123",
selected_organization_id=42,
),
)
assert response == {"status": [{"model": "all", "message": "ok"}]}
validate.assert_awaited_once_with(
stale_config,
organization_id=42,
created_by="provider-123",
)
touch_validation_cache.assert_awaited_once_with(42)
@pytest.mark.asyncio
async def test_validate_user_configurations_uses_fresh_org_v2_validation_cache(
monkeypatch,
):
fresh_config = EffectiveAIModelConfiguration(last_validated_at=datetime.now(UTC))
resolved = ResolvedAIModelConfiguration(
effective=fresh_config,
source="organization_v2",
)
validate = AsyncMock()
touch_validation_cache = AsyncMock()
class FakeValidator:
def __init__(self):
self.validate = validate
monkeypatch.setattr(
user_routes,
"get_resolved_ai_model_configuration",
AsyncMock(return_value=resolved),
)
monkeypatch.setattr(user_routes, "UserConfigurationValidator", FakeValidator)
monkeypatch.setattr(
user_routes,
"update_organization_ai_model_configuration_last_validated_at",
touch_validation_cache,
)
response = await user_routes.validate_user_configurations(
validity_ttl_seconds=60,
user=SimpleNamespace(
provider_id="provider-123",
selected_organization_id=42,
),
)
assert response == {"status": []}
validate.assert_not_awaited()
touch_validation_cache.assert_not_awaited()