mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-09-16 14:57:15 +00:00
feat(provider): support Codex subscriptions with ChatGPT sign-in (#2513)
* feat(provider): support Codex subscriptions with ChatGPT sign-in * style: format Codex live integration test * fix(provider): preserve Codex identity in temporary model tests * fix(web): portal provider selector without dialog overflow * fix(web): allow native scrolling in provider dropdown * fix(provider): surface safe Codex quota and upstream errors * fix(web): provide reliable Codex copy feedback in dialogs * feat(provider): confirm cascade deletion from edit dialog * fix(persistence): discard connections after failed commit * fix(web): polish provider loading and confirmation motion --------- Co-authored-by: dadachann <185672915+dadachann@users.noreply.github.com>
This commit is contained in:
@@ -20,6 +20,7 @@ from langbot.pkg.entity.persistence.model import LLMModel, ModelProvider
|
||||
from langbot.pkg.entity.persistence.pipeline import LegacyPipeline
|
||||
from langbot.pkg.entity.persistence.workspace import Workspace
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
from langbot.pkg.persistence.mgr import PersistenceManager
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
@@ -28,15 +29,10 @@ WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
|
||||
WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
|
||||
|
||||
|
||||
class _PersistenceManager:
|
||||
class _PersistenceManager(PersistenceManager):
|
||||
def __init__(self, engine):
|
||||
self.engine = engine
|
||||
|
||||
async def execute_async(self, *args, **kwargs):
|
||||
async with self.engine.connect() as connection:
|
||||
result = await connection.execute(*args, **kwargs)
|
||||
await connection.commit()
|
||||
return result
|
||||
super().__init__(SimpleNamespace())
|
||||
self.db = SimpleNamespace(get_engine=lambda: engine)
|
||||
|
||||
@staticmethod
|
||||
def serialize_model(model, data, masked_columns=None):
|
||||
|
||||
@@ -0,0 +1,401 @@
|
||||
"""Provider deletion uses real SQLite transactions and real runtime cache cleanup."""
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
import quart
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
from langbot.pkg.api.http.controller.groups.provider.providers import ModelProvidersRouterGroup
|
||||
from langbot.pkg.api.http.service.provider import ModelProviderService
|
||||
from langbot.pkg.entity.persistence.model import CodexCredential, EmbeddingModel, LLMModel, ModelProvider, RerankModel
|
||||
from langbot.pkg.entity.persistence.user import User
|
||||
from langbot.pkg.entity.persistence.workspace import Workspace
|
||||
from langbot.pkg.persistence.mgr import PersistenceManager, PersistenceMode
|
||||
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
MODEL_TYPES = (LLMModel, EmbeddingModel, RerankModel)
|
||||
TABLES = (*MODEL_TYPES, CodexCredential, ModelProvider)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def deletion(tmp_path):
|
||||
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "cascade.db"}')
|
||||
|
||||
@sa.event.listens_for(engine.sync_engine, 'connect')
|
||||
def enable_foreign_keys(connection, _record):
|
||||
connection.execute('PRAGMA foreign_keys=ON')
|
||||
|
||||
ap = SimpleNamespace(logger=Mock())
|
||||
pm = ap.persistence_mgr = PersistenceManager(ap)
|
||||
pm.db = SimpleNamespace(get_engine=lambda: engine)
|
||||
manager = ap.model_mgr = ModelManager(ap)
|
||||
contexts = {workspace: ExecutionContext('instance', workspace, 1) for workspace in ('a', 'b')}
|
||||
# Only execution binding discovery is stubbed; cache indexing/removal/close is real.
|
||||
manager.resolve_execution_context = AsyncMock(side_effect=lambda context: contexts[context])
|
||||
service = ModelProviderService(ap)
|
||||
closed = []
|
||||
|
||||
async def snapshot():
|
||||
async with engine.connect() as conn:
|
||||
return {
|
||||
table.__tablename__: [dict(row) for row in (await conn.execute(sa.select(table))).mappings()]
|
||||
for table in TABLES
|
||||
}
|
||||
|
||||
async with engine.begin() as conn:
|
||||
for table in (User, Workspace, ModelProvider, CodexCredential, *MODEL_TYPES):
|
||||
await conn.run_sync(table.__table__.create)
|
||||
for workspace in contexts:
|
||||
await conn.execute(
|
||||
sa.insert(Workspace).values(
|
||||
uuid=workspace,
|
||||
instance_uuid='instance',
|
||||
name=workspace,
|
||||
slug=workspace,
|
||||
source='cloud_projection',
|
||||
)
|
||||
)
|
||||
for provider, workspace in (('target', 'a'), ('neighbor', 'a'), ('foreign', 'b'), ('empty', 'a')):
|
||||
await conn.execute(
|
||||
sa.insert(ModelProvider).values(
|
||||
uuid=provider,
|
||||
workspace_uuid=workspace,
|
||||
name=provider,
|
||||
requester='openai-codex',
|
||||
base_url='https://chatgpt.com/backend-api/codex',
|
||||
api_keys=[],
|
||||
)
|
||||
)
|
||||
await conn.execute(
|
||||
sa.insert(CodexCredential).values(
|
||||
provider_uuid=provider,
|
||||
workspace_uuid=workspace,
|
||||
payload={'synthetic': provider},
|
||||
)
|
||||
)
|
||||
|
||||
async def close(provider=provider):
|
||||
# A separate connection must observe the durable deletion before close runs.
|
||||
state = await snapshot()
|
||||
assert all(row['uuid'] != provider for row in state['model_providers'])
|
||||
assert pm.current_session() is None
|
||||
closed.append(provider)
|
||||
|
||||
runtime = SimpleNamespace(requester=SimpleNamespace(aclose=AsyncMock(side_effect=close)))
|
||||
manager._cache_set(manager.provider_dict, manager._cache_key(contexts[workspace], provider), runtime)
|
||||
if provider == 'empty':
|
||||
continue
|
||||
for model_type, cache in zip(
|
||||
MODEL_TYPES,
|
||||
(
|
||||
manager.llm_model_dict,
|
||||
manager.embedding_model_dict,
|
||||
manager.rerank_model_dict,
|
||||
),
|
||||
):
|
||||
for index in range(2):
|
||||
uuid = f'{provider}-{model_type.__tablename__}-{index}'
|
||||
await conn.execute(
|
||||
sa.insert(model_type).values(
|
||||
uuid=uuid,
|
||||
workspace_uuid=workspace,
|
||||
provider_uuid=provider,
|
||||
name=uuid,
|
||||
)
|
||||
)
|
||||
manager._cache_set(cache, manager._cache_key(contexts[workspace], uuid), object())
|
||||
initial = await snapshot()
|
||||
initial_caches = [
|
||||
dict(cache)
|
||||
for cache in (
|
||||
manager.provider_dict,
|
||||
manager.llm_model_dict,
|
||||
manager.embedding_model_dict,
|
||||
manager.rerank_model_dict,
|
||||
)
|
||||
]
|
||||
try:
|
||||
yield SimpleNamespace(
|
||||
ap=ap,
|
||||
pm=pm,
|
||||
engine=engine,
|
||||
service=service,
|
||||
manager=manager,
|
||||
snapshot=snapshot,
|
||||
initial=initial,
|
||||
initial_caches=initial_caches,
|
||||
closed=closed,
|
||||
)
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def assert_caches_unchanged(deletion):
|
||||
assert deletion.closed == []
|
||||
assert deletion.initial_caches == [
|
||||
dict(cache)
|
||||
for cache in (
|
||||
deletion.manager.provider_dict,
|
||||
deletion.manager.llm_model_dict,
|
||||
deletion.manager.embedding_model_dict,
|
||||
deletion.manager.rerank_model_dict,
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize('mode', [PersistenceMode.OSS_COMPAT, PersistenceMode.CLOUD_RUNTIME])
|
||||
async def test_cascade_deletes_all_model_types_and_credentials_after_commit(deletion, mode):
|
||||
deletion.pm.mode = mode
|
||||
await deletion.service.delete_provider('a', 'target', cascade=True)
|
||||
state = await deletion.snapshot()
|
||||
for table, rows in deletion.initial.items():
|
||||
identity = 'uuid' if table == 'model_providers' else 'provider_uuid'
|
||||
assert state[table] == [row for row in rows if row[identity] != 'target']
|
||||
assert deletion.closed == ['target']
|
||||
deletion.ap.logger.warning.assert_not_called()
|
||||
for cache in (
|
||||
deletion.manager.provider_dict,
|
||||
deletion.manager.llm_model_dict,
|
||||
deletion.manager.embedding_model_dict,
|
||||
deletion.manager.rerank_model_dict,
|
||||
):
|
||||
assert all(not key[-1].startswith('target') for key in cache)
|
||||
assert any(key[1] == 'b' for key in cache)
|
||||
assert any(key[-1].startswith('neighbor') for key in cache)
|
||||
|
||||
|
||||
@pytest.mark.parametrize('model_type', MODEL_TYPES)
|
||||
@pytest.mark.parametrize('kwargs', [{}, {'cascade': False}])
|
||||
async def test_default_guard_preserves_each_model_type(deletion, model_type, kwargs):
|
||||
async with deletion.engine.begin() as conn:
|
||||
for other in MODEL_TYPES:
|
||||
if other is not model_type:
|
||||
await conn.execute(sa.delete(other).where(other.provider_uuid == 'target'))
|
||||
before = await deletion.snapshot()
|
||||
with pytest.raises(ValueError, match='models still reference it'):
|
||||
await deletion.service.delete_provider('a', 'target', **kwargs)
|
||||
assert await deletion.snapshot() == before
|
||||
assert_caches_unchanged(deletion)
|
||||
|
||||
|
||||
@pytest.mark.parametrize('kwargs', [{}, {'cascade': False}, {'cascade': True}])
|
||||
async def test_empty_provider_deletes_credentials_with_or_without_cascade(deletion, kwargs):
|
||||
await deletion.service.delete_provider('a', 'empty', **kwargs)
|
||||
state = await deletion.snapshot()
|
||||
assert all(row['uuid'] != 'empty' for row in state['model_providers'])
|
||||
assert all(row['provider_uuid'] != 'empty' for row in state['codex_credentials'])
|
||||
assert deletion.closed == ['empty']
|
||||
|
||||
|
||||
@pytest.mark.parametrize('provider', ['foreign', 'missing'])
|
||||
@pytest.mark.parametrize('cascade', [False, True])
|
||||
async def test_foreign_and_missing_provider_are_non_enumerating(deletion, provider, cascade):
|
||||
with pytest.raises(WorkspaceNotFoundError, match='Provider not found'):
|
||||
await deletion.service.delete_provider('a', provider, cascade=cascade)
|
||||
assert await deletion.snapshot() == deletion.initial
|
||||
assert_caches_unchanged(deletion)
|
||||
|
||||
|
||||
@pytest.mark.parametrize('cascade', [False, True])
|
||||
async def test_cloud_managed_provider_cannot_be_deleted(deletion, cascade):
|
||||
async with deletion.engine.begin() as conn:
|
||||
await conn.execute(
|
||||
sa.update(ModelProvider)
|
||||
.where(ModelProvider.uuid == 'target')
|
||||
.values(
|
||||
requester='space-chat-completions',
|
||||
)
|
||||
)
|
||||
before = await deletion.snapshot()
|
||||
deletion.pm.mode = PersistenceMode.CLOUD_RUNTIME
|
||||
with pytest.raises(ValueError, match='managed by Cloud'):
|
||||
await deletion.service.delete_provider('a', 'target', cascade=cascade)
|
||||
assert await deletion.snapshot() == before
|
||||
assert_caches_unchanged(deletion)
|
||||
|
||||
|
||||
@pytest.mark.parametrize('failure_table', ['embedding_models', 'codex_credentials', 'model_providers'])
|
||||
async def test_database_failure_rolls_back_all_rows_without_runtime_cleanup(deletion, failure_table):
|
||||
async with deletion.engine.begin() as conn:
|
||||
await conn.exec_driver_sql(
|
||||
f'CREATE TRIGGER fail_delete BEFORE DELETE ON {failure_table} '
|
||||
"BEGIN SELECT RAISE(ABORT, 'injected delete failure'); END"
|
||||
)
|
||||
with pytest.raises(sa.exc.IntegrityError, match='injected delete failure'):
|
||||
await deletion.service.delete_provider('a', 'target', cascade=True)
|
||||
assert await deletion.snapshot() == deletion.initial
|
||||
assert_caches_unchanged(deletion)
|
||||
|
||||
|
||||
@pytest.mark.parametrize('rollback', [False, True])
|
||||
async def test_nested_transaction_defers_cleanup_until_outer_commit(deletion, rollback):
|
||||
class Abort(Exception):
|
||||
pass
|
||||
|
||||
try:
|
||||
async with deletion.pm.tenant_uow('a'):
|
||||
await deletion.service.delete_provider('a', 'target', cascade=True)
|
||||
assert_caches_unchanged(deletion)
|
||||
if rollback:
|
||||
raise Abort
|
||||
except Abort:
|
||||
pass
|
||||
tasks = tuple(deletion.service._deletion_tasks)
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
if rollback:
|
||||
assert await deletion.snapshot() == deletion.initial
|
||||
assert_caches_unchanged(deletion)
|
||||
else:
|
||||
assert deletion.closed == ['target']
|
||||
deletion.ap.logger.warning.assert_not_called()
|
||||
|
||||
|
||||
async def test_cascade_ignores_foreign_workspace_references_even_without_foreign_keys(deletion):
|
||||
async with deletion.engine.connect() as conn:
|
||||
await conn.exec_driver_sql('PRAGMA foreign_keys=OFF')
|
||||
for model_type in MODEL_TYPES:
|
||||
await conn.execute(
|
||||
sa.update(model_type)
|
||||
.where(model_type.workspace_uuid == 'b')
|
||||
.values(
|
||||
provider_uuid='target',
|
||||
)
|
||||
)
|
||||
await conn.commit()
|
||||
await deletion.service.delete_provider('a', 'target', cascade=True)
|
||||
state = await deletion.snapshot()
|
||||
for model_type in MODEL_TYPES:
|
||||
assert len([row for row in state[model_type.__tablename__] if row['workspace_uuid'] == 'b']) == 2
|
||||
assert all(row['provider_uuid'] != 'target' for row in state['codex_credentials'])
|
||||
assert deletion.closed == ['target']
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def route_app(deletion):
|
||||
ap = deletion.ap
|
||||
ap.user_service = SimpleNamespace(
|
||||
get_authenticated_account=AsyncMock(
|
||||
return_value=SimpleNamespace(uuid='account', user='owner@example.invalid'),
|
||||
)
|
||||
)
|
||||
membership = SimpleNamespace(uuid='membership', role='owner', projection_revision=0)
|
||||
ap.workspace_collaboration_service = SimpleNamespace(
|
||||
resolve_account_workspace=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
workspace=SimpleNamespace(uuid='a'),
|
||||
membership=membership,
|
||||
execution=SimpleNamespace(instance_uuid='instance', placement_generation=1),
|
||||
),
|
||||
)
|
||||
)
|
||||
ap.provider_service = SimpleNamespace(delete_provider=AsyncMock())
|
||||
app = quart.Quart(__name__)
|
||||
await ModelProvidersRouterGroup(ap, app).initialize()
|
||||
return app.test_client(), ap.provider_service.delete_provider, membership
|
||||
|
||||
|
||||
@pytest.mark.parametrize('query, expected', [('', None), ('?cascade=true', True), ('?cascade=false', False)])
|
||||
async def test_route_passes_explicit_cascade_and_trusted_workspace(route_app, query, expected):
|
||||
client, delete, _ = route_app
|
||||
response = await client.delete(
|
||||
'/api/v1/provider/providers/target' + query,
|
||||
headers={
|
||||
'Authorization': 'Bearer token',
|
||||
'X-Workspace-Id': 'a',
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
assert delete.await_count == 1
|
||||
assert delete.await_args.args[0].workspace_uuid == 'a'
|
||||
assert delete.await_args.args[1] == 'target'
|
||||
assert delete.await_args.kwargs == ({} if expected is None else {'cascade': expected})
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
'query',
|
||||
[
|
||||
'?cascade=',
|
||||
'?cascade',
|
||||
'?cascade=TRUE',
|
||||
'?cascade=1',
|
||||
'?cascade=yes',
|
||||
'?cascade=null',
|
||||
'?cascade=%20true',
|
||||
'?cascade=true&cascade=false',
|
||||
'?cascade=true&cascade=true',
|
||||
],
|
||||
)
|
||||
async def test_route_rejects_invalid_or_duplicate_cascade_before_deletion(route_app, query):
|
||||
client, delete, _ = route_app
|
||||
response = await client.delete(
|
||||
'/api/v1/provider/providers/target' + query,
|
||||
headers={
|
||||
'Authorization': 'Bearer token',
|
||||
'X-Workspace-Id': 'a',
|
||||
},
|
||||
)
|
||||
assert response.status_code == 400
|
||||
delete.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize('role', ['viewer', 'operator'])
|
||||
async def test_cascade_requires_workspace_resource_manage_permission(route_app, role):
|
||||
client, delete, membership = route_app
|
||||
membership.role = role
|
||||
response = await client.delete(
|
||||
'/api/v1/provider/providers/target?cascade=true',
|
||||
headers={
|
||||
'Authorization': 'Bearer token',
|
||||
'X-Workspace-Id': 'a',
|
||||
},
|
||||
)
|
||||
assert response.status_code == 403
|
||||
delete.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
'provider, query, status',
|
||||
[
|
||||
('target', '', 400),
|
||||
('target', '?cascade=false', 400),
|
||||
('target', '?cascade=true', 200),
|
||||
('foreign', '?cascade=true', 404),
|
||||
('missing', '?cascade=true', 404),
|
||||
],
|
||||
)
|
||||
async def test_route_to_real_sqlite_service(deletion, route_app, provider, query, status):
|
||||
client, _, _ = route_app
|
||||
deletion.ap.provider_service = deletion.service
|
||||
# The route forwards RequestContext, unlike the string-context service tests.
|
||||
deletion.manager.resolve_execution_context = AsyncMock(
|
||||
side_effect=lambda context: ExecutionContext(
|
||||
context.instance_uuid,
|
||||
context.workspace_uuid,
|
||||
context.placement_generation,
|
||||
)
|
||||
)
|
||||
response = await client.delete(
|
||||
'/api/v1/provider/providers/' + provider + query,
|
||||
headers={
|
||||
'Authorization': 'Bearer token',
|
||||
'X-Workspace-Id': 'a',
|
||||
},
|
||||
)
|
||||
assert response.status_code == status
|
||||
if status == 200:
|
||||
assert deletion.closed == ['target']
|
||||
for model_type in MODEL_TYPES:
|
||||
assert all(
|
||||
row['provider_uuid'] != 'target' for row in (await deletion.snapshot())[model_type.__tablename__]
|
||||
)
|
||||
else:
|
||||
assert await deletion.snapshot() == deletion.initial
|
||||
assert_caches_unchanged(deletion)
|
||||
@@ -14,11 +14,12 @@ Source: src/langbot/pkg/api/http/service/provider.py
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from contextlib import nullcontext
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
from types import SimpleNamespace
|
||||
|
||||
from langbot.pkg.api.http.service.provider import ModelProviderService
|
||||
from langbot.pkg.entity.persistence.model import ModelProvider, LLMModel, EmbeddingModel, RerankModel
|
||||
from langbot.pkg.entity.persistence.model import ModelProvider, LLMModel
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
|
||||
|
||||
@@ -383,112 +384,35 @@ class TestModelProviderServiceUpdateProvider:
|
||||
|
||||
|
||||
class TestModelProviderServiceDeleteProvider:
|
||||
"""Tests for delete_provider method."""
|
||||
|
||||
async def test_delete_provider_with_llm_models_raises_error(self):
|
||||
"""Raises ValueError when LLM models reference provider."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
|
||||
# Mock LLM model exists - only return LLM result since that's first check
|
||||
llm_result = _create_mock_result([], first_item=_create_mock_llm_model())
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=llm_result)
|
||||
"""Fast guard coverage; real transaction/cache behavior is in test_provider_cascade."""
|
||||
|
||||
@pytest.mark.parametrize('label', ['LLM', 'Embedding', 'Rerank', None])
|
||||
async def test_delete_provider_requires_no_references(self, label):
|
||||
provider_result = Mock()
|
||||
provider_result.first.return_value = SimpleNamespace(requester='openai')
|
||||
results = [provider_result]
|
||||
for model_label in ('LLM', 'Embedding', 'Rerank'):
|
||||
result = Mock()
|
||||
result.scalars.return_value = ['model'] if label == model_label else []
|
||||
results.append(result)
|
||||
results.extend([Mock(rowcount=1), Mock(rowcount=1)])
|
||||
ap = SimpleNamespace(
|
||||
persistence_mgr=SimpleNamespace(
|
||||
execute_async=AsyncMock(side_effect=results),
|
||||
tenant_uow=lambda _: nullcontext(),
|
||||
tenant_scope=lambda _: nullcontext(),
|
||||
current_session=lambda: None,
|
||||
),
|
||||
model_mgr=SimpleNamespace(remove_provider=AsyncMock()),
|
||||
)
|
||||
service = ModelProviderService(ap)
|
||||
|
||||
# Execute & Verify
|
||||
with pytest.raises(ValueError, match='Cannot delete provider: LLM models'):
|
||||
await service.delete_provider(WORKSPACE_UUID, 'provider-with-llm')
|
||||
|
||||
async def test_delete_provider_with_embedding_models_raises_error(self):
|
||||
"""Raises ValueError when Embedding models reference provider."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
|
||||
# Create results for each check type
|
||||
llm_result = Mock()
|
||||
llm_result.first = Mock(return_value=None) # No LLM models
|
||||
embedding_result = Mock()
|
||||
embedding_result.first = Mock(return_value=Mock(spec=EmbeddingModel)) # Has embedding model
|
||||
rerank_result = Mock()
|
||||
rerank_result.first = Mock(return_value=None)
|
||||
|
||||
call_count = 0
|
||||
|
||||
async def mock_execute(query):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
return llm_result
|
||||
elif call_count == 2:
|
||||
return embedding_result
|
||||
return rerank_result
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
|
||||
service = ModelProviderService(ap)
|
||||
|
||||
# Execute & Verify - should raise embedding error (LLM check passes, embedding check fails)
|
||||
with pytest.raises(ValueError, match='Cannot delete provider: Embedding models'):
|
||||
await service.delete_provider(WORKSPACE_UUID, 'provider-with-embedding')
|
||||
|
||||
async def test_delete_provider_with_rerank_models_raises_error(self):
|
||||
"""Raises ValueError when Rerank models reference provider."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
|
||||
# Create results for each check type
|
||||
llm_result = Mock()
|
||||
llm_result.first = Mock(return_value=None) # No LLM models
|
||||
embedding_result = Mock()
|
||||
embedding_result.first = Mock(return_value=None) # No embedding models
|
||||
rerank_result = Mock()
|
||||
rerank_result.first = Mock(return_value=Mock(spec=RerankModel)) # Has rerank model
|
||||
|
||||
call_count = 0
|
||||
|
||||
async def mock_execute(query):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
return llm_result
|
||||
elif call_count == 2:
|
||||
return embedding_result
|
||||
return rerank_result
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
|
||||
service = ModelProviderService(ap)
|
||||
|
||||
# Execute & Verify - should raise rerank error (LLM and embedding checks pass, rerank check fails)
|
||||
with pytest.raises(ValueError, match='Cannot delete provider: Rerank models'):
|
||||
await service.delete_provider(WORKSPACE_UUID, 'provider-with-rerank')
|
||||
|
||||
async def test_delete_provider_no_models_success(self):
|
||||
"""Deletes provider when no models reference it."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.model_mgr = SimpleNamespace()
|
||||
ap.model_mgr.remove_provider = AsyncMock()
|
||||
|
||||
# Mock no models reference provider
|
||||
empty_result = Mock()
|
||||
empty_result.first = Mock(return_value=None)
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=empty_result)
|
||||
|
||||
service = ModelProviderService(ap)
|
||||
|
||||
# Execute
|
||||
await service.delete_provider(WORKSPACE_UUID, 'provider-no-models')
|
||||
|
||||
# Verify - delete and remove called
|
||||
ap.model_mgr.remove_provider.assert_called_once_with(WORKSPACE_UUID, 'provider-no-models')
|
||||
if label is not None:
|
||||
with pytest.raises(ValueError, match=f'Cannot delete provider: {label} models'):
|
||||
await service.delete_provider(WORKSPACE_UUID, 'provider')
|
||||
ap.model_mgr.remove_provider.assert_not_awaited()
|
||||
else:
|
||||
await service.delete_provider(WORKSPACE_UUID, 'provider')
|
||||
ap.model_mgr.remove_provider.assert_awaited_once_with(WORKSPACE_UUID, 'provider')
|
||||
|
||||
|
||||
class TestModelProviderServiceGetProviderModelCounts:
|
||||
@@ -1045,15 +969,18 @@ class TestCloudManagedProviderProtection:
|
||||
|
||||
async def test_cloud_rejects_update_and_delete_of_managed_provider(self):
|
||||
service = self._service()
|
||||
service.get_provider = AsyncMock(
|
||||
return_value={'uuid': 'system-provider', 'requester': SYSTEM_REQUESTER}
|
||||
)
|
||||
service.get_provider = AsyncMock(return_value={'uuid': 'system-provider', 'requester': SYSTEM_REQUESTER})
|
||||
|
||||
with pytest.raises(ValueError, match='managed by Cloud'):
|
||||
await service.update_provider(WORKSPACE_UUID, 'system-provider', {'name': 'Renamed'})
|
||||
service.ap.persistence_mgr.execute_async.assert_not_awaited()
|
||||
service.ap.persistence_mgr.tenant_uow = lambda _: nullcontext()
|
||||
result = Mock()
|
||||
result.first.return_value = SimpleNamespace(requester=SYSTEM_REQUESTER)
|
||||
service.ap.persistence_mgr.execute_async.return_value = result
|
||||
with pytest.raises(ValueError, match='managed by Cloud'):
|
||||
await service.delete_provider(WORKSPACE_UUID, 'system-provider')
|
||||
service.ap.persistence_mgr.execute_async.assert_not_awaited()
|
||||
assert service.ap.persistence_mgr.execute_async.await_count == 1
|
||||
|
||||
async def test_oss_does_not_reserve_space_requester(self):
|
||||
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(mode=SimpleNamespace(value='oss_compat')))
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
"""Keep subprocess coverage pointed at the generated E2E configuration."""
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from tests.e2e.utils.process_manager import LangBotProcess
|
||||
|
||||
|
||||
def test_e2e_coverage_environment_uses_generated_config(tmp_path):
|
||||
process = Mock()
|
||||
process.poll.return_value = None
|
||||
project = tmp_path / 'project'
|
||||
project.mkdir()
|
||||
manager = LangBotProcess(project, tmp_path, collect_coverage=True)
|
||||
with patch('subprocess.Popen', return_value=process) as popen, patch('httpx.get') as get:
|
||||
get.return_value.status_code = 200
|
||||
assert manager.start()
|
||||
config = Path(popen.call_args.kwargs['env']['COVERAGE_PROCESS_START'])
|
||||
assert config.is_file()
|
||||
assert f'--rcfile={config}' in popen.call_args.args[0]
|
||||
@@ -0,0 +1,115 @@
|
||||
"""Real rollback-journal contention must not poison the pooled writer."""
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import sqlite3
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.ext.asyncio import AsyncConnection, AsyncSessionTransaction, create_async_engine
|
||||
|
||||
from langbot.pkg.persistence.mgr import PersistenceManager, PersistenceMode
|
||||
from langbot.pkg.persistence.tenant_uow import TenantScopedAsyncSession
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('cancel_commit', [False, True])
|
||||
@pytest.mark.parametrize('close_fails', [False, True])
|
||||
@pytest.mark.parametrize('cancel_cleanup', [False, True])
|
||||
async def test_failed_commit_releases_sqlite_writer_and_scope(
|
||||
tmp_path, monkeypatch, cancel_commit, close_fails, cancel_cleanup
|
||||
):
|
||||
path = tmp_path / 'failed-commit.db'
|
||||
engine = create_async_engine(
|
||||
f'sqlite+aiosqlite:///{path}', connect_args={'timeout': 0.05}, pool_size=1, max_overflow=0
|
||||
)
|
||||
table = sa.Table('rows', sa.MetaData(), sa.Column('id', sa.Integer, primary_key=True))
|
||||
manager = PersistenceManager(object(), mode=PersistenceMode.CLOUD_RUNTIME)
|
||||
manager.db = SimpleNamespace(get_engine=lambda: engine)
|
||||
original_commit = AsyncSessionTransaction.commit
|
||||
original_error = None
|
||||
invalidation_finished = False
|
||||
owner = asyncio.current_task()
|
||||
original_invalidate = AsyncConnection.invalidate
|
||||
|
||||
async def delayed_invalidate(connection, exception=None):
|
||||
nonlocal invalidation_finished
|
||||
if cancel_cleanup:
|
||||
owner.cancel()
|
||||
await asyncio.sleep(0)
|
||||
owner.cancel()
|
||||
await asyncio.sleep(0)
|
||||
await original_invalidate(connection, exception)
|
||||
invalidation_finished = True
|
||||
|
||||
monkeypatch.setattr(AsyncConnection, 'invalidate', delayed_invalidate)
|
||||
|
||||
async def failing_commit(transaction):
|
||||
nonlocal original_error
|
||||
try:
|
||||
await original_commit(transaction)
|
||||
except sa.exc.OperationalError as exc:
|
||||
original_error = asyncio.CancelledError('commit cancelled') if cancel_commit else exc
|
||||
raise original_error
|
||||
|
||||
monkeypatch.setattr(AsyncSessionTransaction, 'commit', failing_commit)
|
||||
original_close = TenantScopedAsyncSession._close_owned_session
|
||||
|
||||
async def failing_close(session, capability):
|
||||
await original_close(session, capability)
|
||||
if original_error is not None:
|
||||
raise RuntimeError('secondary close failure')
|
||||
|
||||
if close_fails:
|
||||
monkeypatch.setattr(TenantScopedAsyncSession, '_close_owned_session', failing_close)
|
||||
blocker = sqlite3.connect(path, timeout=0.05)
|
||||
try:
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(table.metadata.create_all)
|
||||
await connection.execute(sa.insert(table).values(id=1))
|
||||
blocker.execute('BEGIN')
|
||||
blocker.execute('SELECT * FROM rows').fetchall()
|
||||
error_type = asyncio.CancelledError if cancel_commit else sa.exc.OperationalError
|
||||
with pytest.raises(error_type) as caught:
|
||||
async with manager.tenant_uow('workspace-a') as outer:
|
||||
gate = manager.create_after_commit_gate()
|
||||
state = outer._active_state
|
||||
async with manager.tenant_uow('workspace-a') as inner:
|
||||
assert inner.session is outer.session
|
||||
await manager.execute_async(sa.insert(table).values(id=2))
|
||||
assert not gate.done()
|
||||
assert state.depth == 1
|
||||
assert caught.value is original_error
|
||||
assert invalidation_finished
|
||||
if cancel_cleanup:
|
||||
# Do not leak the synthetic cancellation count into pytest.
|
||||
owner.uncancel()
|
||||
owner.uncancel()
|
||||
monkeypatch.setattr(TenantScopedAsyncSession, '_close_owned_session', original_close)
|
||||
if close_fails:
|
||||
assert any('secondary close failure' in note for note in caught.value.__notes__)
|
||||
assert gate.cancelled()
|
||||
assert state.depth == 0
|
||||
assert manager.current_session() is None
|
||||
with pytest.raises(RuntimeError, match='not active'):
|
||||
_ = outer.session
|
||||
|
||||
# The original SHARED lock remains. New reads and RESERVED writes
|
||||
# must work; COMMIT of another write must wait for its release.
|
||||
assert blocker.in_transaction
|
||||
with contextlib.closing(sqlite3.connect(path, timeout=0.05)) as probe:
|
||||
assert probe.execute('SELECT id FROM rows').fetchall() == [(1,)]
|
||||
probe.execute('INSERT INTO rows VALUES (3)')
|
||||
probe.rollback()
|
||||
async with manager.tenant_uow('workspace-b'):
|
||||
assert (await manager.execute_async(sa.select(table.c.id))).scalars().all() == [1]
|
||||
blocker.rollback()
|
||||
async with manager.tenant_uow('workspace-b'):
|
||||
await manager.execute_async(sa.insert(table).values(id=4))
|
||||
async with engine.connect() as connection:
|
||||
assert (await connection.execute(sa.select(table.c.id).order_by(table.c.id))).scalars().all() == [1, 4]
|
||||
assert engine.pool.checkedout() == 0
|
||||
finally:
|
||||
blocker.close()
|
||||
await engine.dispose()
|
||||
@@ -0,0 +1,200 @@
|
||||
"""Replay synthetic HTTP/SSE traffic through the real Codex requester."""
|
||||
|
||||
import json
|
||||
from collections import OrderedDict
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import langbot_plugin.api.entities.builtin.provider.message as pm
|
||||
|
||||
from langbot.pkg.provider.modelmgr.requesters.codex import CodexRequester, sse_events
|
||||
|
||||
|
||||
TOKENS = {'access_token': 'access-secret', 'account_id': 'account', 'connection_id': 'connection'}
|
||||
MODEL = SimpleNamespace(model_entity=SimpleNamespace(name='codex-test', extra_args={}, reasoning_config=None))
|
||||
|
||||
|
||||
def requester(monkeypatch, handler):
|
||||
real_client = httpx.AsyncClient
|
||||
monkeypatch.setattr(
|
||||
httpx, 'AsyncClient', lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs)
|
||||
)
|
||||
obj = object.__new__(CodexRequester)
|
||||
obj.workspace, obj.provider = 'w', 'p'
|
||||
obj._replay = OrderedDict()
|
||||
obj.auth = SimpleNamespace(access=AsyncMock(return_value=TOKENS))
|
||||
return obj
|
||||
|
||||
|
||||
def stream(events):
|
||||
return httpx.Response(200, content=''.join('data: ' + json.dumps(event) + '\r\n\r\n' for event in events))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_tools_usage_and_scoped_opaque_replay(monkeypatch):
|
||||
call = {'type': 'function_call', 'call_id': 'call_1', 'name': 'lookup', 'arguments': '{"q":"test"}'}
|
||||
output = [
|
||||
{'type': 'reasoning', 'encrypted_content': 'opaque-secret'},
|
||||
call,
|
||||
{'type': 'message', 'role': 'assistant', 'content': [{'type': 'output_text', 'text': 'Hello'}]},
|
||||
]
|
||||
requests = []
|
||||
|
||||
def handler(request):
|
||||
requests.append(request)
|
||||
return stream(
|
||||
[
|
||||
{'type': 'response.created', 'response': {'id': 'resp_1'}},
|
||||
{'type': 'response.output_text.delta', 'delta': 'Hel'},
|
||||
{'type': 'response.output_text.delta', 'delta': 'lo'},
|
||||
{'type': 'response.function_call_arguments.delta', 'delta': '{broken'},
|
||||
{'type': 'response.output_item.done', 'item': call, 'output_index': 1},
|
||||
{
|
||||
'type': 'response.completed',
|
||||
'response': {
|
||||
'id': 'resp_1',
|
||||
'status': 'completed',
|
||||
'output': output,
|
||||
'usage': {'input_tokens': 4, 'output_tokens': 3, 'input_tokens_details': {'cached_tokens': 2}},
|
||||
},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
obj = requester(monkeypatch, handler)
|
||||
query = SimpleNamespace(query_id='q', variables=None)
|
||||
messages = [pm.Message(role='system', content='Be brief'), pm.Message(role='user', content='Hi')]
|
||||
message, usage = await obj.invoke_llm(query, MODEL, messages)
|
||||
assert message.content == 'Hello'
|
||||
assert len(message.tool_calls) == 1
|
||||
assert message.tool_calls[0].function.arguments == '{"q":"test"}'
|
||||
assert usage['total_tokens'] == 7
|
||||
assert query.variables['_stream_usage'] == usage
|
||||
assert 'opaque-secret' not in message.model_dump_json()
|
||||
body = json.loads(requests[0].content)
|
||||
assert body['store'] is False and body['stream'] is True
|
||||
assert body['instructions'] == 'Be brief'
|
||||
assert requests[0].url.path.endswith('/codex/responses')
|
||||
assert requests[0].headers['authorization'] == 'Bearer access-secret'
|
||||
assert requests[0].headers['originator'] == 'langbot'
|
||||
same = obj._body(query, MODEL, [message], None, None, TOKENS)
|
||||
assert same['input'] == output
|
||||
other = obj._body(SimpleNamespace(query_id='q'), MODEL, [message], None, None, TOKENS)
|
||||
assert 'opaque-secret' not in json.dumps(other)
|
||||
rotated = obj._body(query, MODEL, [message], None, None, {**TOKENS, 'connection_id': 'new'})
|
||||
assert 'opaque-secret' not in json.dumps(rotated)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
'events',
|
||||
[
|
||||
[{'type': 'response.failed', 'error': 'access-secret'}],
|
||||
[{'type': 'response.incomplete'}],
|
||||
[{'type': 'error'}],
|
||||
[{'type': 'response.output_text.delta', 'delta': 'partial'}],
|
||||
[],
|
||||
],
|
||||
)
|
||||
async def test_failure_and_truncated_stream_never_succeed(monkeypatch, events):
|
||||
obj = requester(monkeypatch, lambda request: stream(events))
|
||||
with pytest.raises(ValueError) as caught:
|
||||
await obj.invoke_llm(None, MODEL, [])
|
||||
assert 'access-secret' not in str(caught.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminal_only_text_stream_and_usage(monkeypatch):
|
||||
obj = requester(
|
||||
monkeypatch,
|
||||
lambda request: stream(
|
||||
[
|
||||
{
|
||||
'type': 'response.done',
|
||||
'response': {
|
||||
'output': [{'type': 'message', 'content': [{'type': 'output_text', 'text': 'done'}]}],
|
||||
'usage': {'input_tokens': 2, 'output_tokens': 1},
|
||||
},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
query = SimpleNamespace(query_id='q', variables={})
|
||||
chunks = [chunk async for chunk in obj.invoke_llm_stream(query, MODEL, [])]
|
||||
assert ''.join(chunk.content or '' for chunk in chunks) == 'done'
|
||||
assert chunks[-1].is_final
|
||||
assert query.variables['_stream_usage']['total_tokens'] == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('status', [401, 403, 429, 500])
|
||||
async def test_http_error_secrecy_and_bounded_401_retry(monkeypatch, status):
|
||||
requests = []
|
||||
|
||||
def handler(request):
|
||||
requests.append(request)
|
||||
return httpx.Response(status, text='access-secret refresh-secret')
|
||||
|
||||
obj = requester(monkeypatch, handler)
|
||||
with pytest.raises(ValueError) as caught:
|
||||
await obj.invoke_llm(None, MODEL, [])
|
||||
assert 'secret' not in str(caught.value)
|
||||
assert len(requests) == (2 if status == 401 else 1)
|
||||
if status == 401:
|
||||
assert obj.auth.access.call_args.kwargs == {'rejected_token': 'access-secret'}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_mapping_filtering_and_deduplication(monkeypatch):
|
||||
requests = []
|
||||
|
||||
def handler(request):
|
||||
requests.append(request)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
'models': [
|
||||
{'slug': 'a', 'input_modalities': ['text', 'image'], 'supported_reasoning_levels': ['low']},
|
||||
{'slug': 'hidden', 'visibility': 'hide'},
|
||||
{'slug': 'a', 'input_modalities': ['text', 'image'], 'supported_reasoning_levels': ['low']},
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
obj = requester(monkeypatch, handler)
|
||||
catalog = await obj.scan_models()
|
||||
assert len(catalog['models']) == 1
|
||||
assert catalog['models'][0]['abilities'] == ['func_call', 'vision', 'reasoning']
|
||||
assert catalog['debug'] is None
|
||||
assert requests[0].url.path.endswith('/codex/models')
|
||||
assert 'client_version' in requests[0].url.params
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('payload', [{'models': ['secret']}, [], {'models': None}])
|
||||
async def test_catalog_malformed_safe(monkeypatch, payload):
|
||||
obj = requester(monkeypatch, lambda request: httpx.Response(200, json=payload))
|
||||
with pytest.raises(ValueError, match='invalid model catalog'):
|
||||
await obj.scan_models()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_multiline_crlf_comments_and_chunk_boundaries():
|
||||
class Bytes(httpx.AsyncByteStream):
|
||||
async def __aiter__(self):
|
||||
for value in b': comment\r\nevent: test\r\ndata: {"type":\r\ndata: "test"}\r\n\r\ndata: [DONE]\r\n\r\n':
|
||||
yield bytes([value])
|
||||
|
||||
response = httpx.Response(200, stream=Bytes())
|
||||
assert [event async for event in sse_events(response)] == [{'type': 'test'}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
'key', ['base_url', 'headers', 'api_key', 'store', 'stream', 'previous_response_id', 'temperature']
|
||||
)
|
||||
def test_advanced_parameters_cannot_override_transport(monkeypatch, key):
|
||||
obj = requester(monkeypatch, lambda request: httpx.Response(200))
|
||||
with pytest.raises(ValueError, match='Unsupported Codex advanced'):
|
||||
obj._body(None, MODEL, [], None, {key: 'secret'}, TOKENS)
|
||||
@@ -0,0 +1,100 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from quart import Quart
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
|
||||
from langbot.pkg.api.http.controller.groups.provider.models import LLMModelsRouterGroup
|
||||
from langbot.pkg.api.http.authz import Permission
|
||||
from langbot.pkg.provider.modelmgr.requesters.codex import CodexRequester
|
||||
from tests.unit_tests.provider.test_codex import requester, MODEL, stream
|
||||
|
||||
|
||||
CASES = [
|
||||
(400, 400, 'codex_invalid_request'),
|
||||
(401, 400, 'codex_reauthentication_required'),
|
||||
(403, 403, 'codex_access_denied'),
|
||||
(429, 429, 'codex_rate_limited'),
|
||||
(500, 502, 'codex_upstream_failure'),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('upstream,status,code', CASES)
|
||||
async def test_requester_safe_error(monkeypatch, upstream, status, code):
|
||||
obj = requester(monkeypatch, lambda request: httpx.Response(upstream, text='credential-secret'))
|
||||
with pytest.raises(Exception) as caught:
|
||||
await obj.invoke_llm(None, MODEL, [])
|
||||
error = caught.value
|
||||
assert getattr(error, 'status_code', None) == status
|
||||
assert error.error_code == code
|
||||
assert 'secret' not in str(error)
|
||||
if upstream == 429:
|
||||
assert 'rate limit' in str(error).lower()
|
||||
assert 'usage limit reached' not in str(error).lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('events', [[{'type': 'response.failed', 'error': 'credential-secret'}], []])
|
||||
async def test_stream_safe_error(monkeypatch, events):
|
||||
obj = requester(monkeypatch, lambda request: stream(events))
|
||||
with pytest.raises(Exception) as caught:
|
||||
await obj.invoke_llm(None, MODEL, [])
|
||||
assert getattr(caught.value, 'status_code', None) == 502
|
||||
assert caught.value.error_code == 'codex_upstream_failure'
|
||||
assert 'secret' not in str(caught.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
'kind,code', [('usage_limit_reached', 'codex_usage_limit_reached'), ('unknown', 'codex_rate_limited')]
|
||||
)
|
||||
async def test_allowlisted_usage_error(monkeypatch, kind, code):
|
||||
obj = requester(
|
||||
monkeypatch,
|
||||
lambda request: httpx.Response(
|
||||
429, json={'error': {'type': kind, 'message': 'credential-secret', 'resets_at': 1789043289}}
|
||||
),
|
||||
)
|
||||
with pytest.raises(Exception) as caught:
|
||||
await obj.invoke_llm(None, MODEL, [])
|
||||
assert caught.value.error_code == code
|
||||
assert 'secret' not in str(caught.value)
|
||||
|
||||
|
||||
async def client_for(error):
|
||||
app = Quart(__name__)
|
||||
ap = SimpleNamespace(logger=Mock(), llm_model_service=SimpleNamespace(test_llm_model=AsyncMock(side_effect=error)))
|
||||
router = LLMModelsRouterGroup(ap, app)
|
||||
router._authenticate_api_key = AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
workspace_uuid='w',
|
||||
workspace=SimpleNamespace(permissions=frozenset({Permission.PROVIDER_SECRET_MANAGE.value})),
|
||||
)
|
||||
)
|
||||
await router.initialize()
|
||||
return app.test_client(), ap
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('upstream,status,code', CASES)
|
||||
async def test_real_model_test_route_safe_error(upstream, status, code):
|
||||
error = CodexRequester._http_error(upstream)
|
||||
client, ap = await client_for(error)
|
||||
response = await client.post('/api/v1/provider/models/llm/model/test', json={}, headers={'X-API-Key': 'synthetic'})
|
||||
body = await response.get_json()
|
||||
assert response.status_code == status
|
||||
assert body['code'] == code
|
||||
assert body['msg'] == str(error)
|
||||
ap.logger.error.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('error', [ValueError('private-value-secret'), SQLAlchemyError('private-sql-secret')])
|
||||
async def test_real_model_test_route_unexpected_errors_hidden(error):
|
||||
client, _ = await client_for(error)
|
||||
response = await client.post('/api/v1/provider/models/llm/model/test', json={}, headers={'X-API-Key': 'synthetic'})
|
||||
assert response.status_code == 500
|
||||
assert 'secret' not in await response.get_data(as_text=True)
|
||||
@@ -0,0 +1,147 @@
|
||||
"""Temporary Codex models use saved, tenant-scoped providers and synthetic SQLite credentials."""
|
||||
|
||||
import time
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from langbot.pkg.entity.persistence.model import CodexCredential, ModelProvider
|
||||
from langbot.pkg.provider.modelmgr.codex_auth import BASE_URL
|
||||
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
|
||||
from langbot.pkg.provider.modelmgr.requesters.codex import CodexRequester
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
from tests.unit_tests.provider.conftest import (
|
||||
TEST_EXECUTION_CONTEXT,
|
||||
TEST_WORKSPACE_UUID,
|
||||
FakeProviderAPIRequester,
|
||||
)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def manager(tmp_path, mock_app_for_modelmgr):
|
||||
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "temporary-codex.db"}')
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(ModelProvider.__table__.create)
|
||||
await conn.run_sync(CodexCredential.__table__.create)
|
||||
for uuid, workspace, kind in (
|
||||
('saved', TEST_WORKSPACE_UUID, 'openai-codex'),
|
||||
('foreign', 'another-workspace', 'openai-codex'),
|
||||
('api', TEST_WORKSPACE_UUID, 'fake-requester'),
|
||||
):
|
||||
await conn.execute(
|
||||
sa.insert(ModelProvider).values(
|
||||
uuid=uuid,
|
||||
workspace_uuid=workspace,
|
||||
name='Saved provider',
|
||||
requester=kind,
|
||||
base_url=BASE_URL,
|
||||
api_keys=[],
|
||||
)
|
||||
)
|
||||
await conn.execute(
|
||||
sa.insert(CodexCredential).values(
|
||||
provider_uuid='saved',
|
||||
workspace_uuid=TEST_WORKSPACE_UUID,
|
||||
payload={
|
||||
'tokens': {
|
||||
'access_token': 'synthetic-access',
|
||||
'refresh_token': 'synthetic-refresh',
|
||||
'account_id': 'synthetic-account',
|
||||
'expires_at': time.time() + 3600,
|
||||
}
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
async def execute(statement):
|
||||
async with engine.begin() as conn:
|
||||
return await conn.execute(statement)
|
||||
|
||||
mock_app_for_modelmgr.persistence_mgr = SimpleNamespace(execute_async=execute)
|
||||
mgr = ModelManager(mock_app_for_modelmgr)
|
||||
mgr.requester_dict = {'openai-codex': CodexRequester, 'fake-requester': FakeProviderAPIRequester}
|
||||
try:
|
||||
yield mgr
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def info(provider_uuid='saved', **inline):
|
||||
result = {'name': 'codex-test', 'provider': {'requester': 'openai-codex', **inline}}
|
||||
if provider_uuid is not None:
|
||||
result['provider_uuid'] = provider_uuid
|
||||
return result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
'inline',
|
||||
[
|
||||
{},
|
||||
{'uuid': 'saved'},
|
||||
{
|
||||
'requester': 'fake-requester',
|
||||
'api_keys': ['untrusted'],
|
||||
'base_url': 'https://untrusted.invalid',
|
||||
'workspace_uuid': 'another-workspace',
|
||||
},
|
||||
],
|
||||
)
|
||||
async def test_codex_temporary_model_resolves_saved_provider_and_real_credentials(manager, inline):
|
||||
model = await manager.init_temporary_runtime_llm_model(TEST_EXECUTION_CONTEXT, info(**inline))
|
||||
provider = model.provider
|
||||
assert provider.provider_entity.uuid == 'saved'
|
||||
assert provider.provider_entity.requester == 'openai-codex'
|
||||
assert provider.provider_entity.api_keys == []
|
||||
assert provider.provider_entity.base_url == BASE_URL
|
||||
assert isinstance(provider.requester, CodexRequester)
|
||||
tokens = await provider.requester.auth.access(provider.requester.workspace, provider.requester.provider)
|
||||
assert tokens['access_token'] == 'synthetic-access'
|
||||
assert model.model_entity.provider_uuid == 'saved'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_temporary_model_accepts_inline_saved_identity(manager):
|
||||
model = await manager.init_temporary_runtime_llm_model(TEST_EXECUTION_CONTEXT, info(None, uuid='saved'))
|
||||
assert model.provider.provider_entity.name == 'Saved provider'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('identity', ['missing', 'foreign', None])
|
||||
async def test_codex_temporary_model_rejects_unavailable_identity(manager, identity):
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await manager.init_temporary_runtime_llm_model(
|
||||
TEST_EXECUTION_CONTEXT, info(identity, workspace_uuid='another-workspace')
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_temporary_model_rejects_non_codex_saved_provider(manager):
|
||||
with pytest.raises(ValueError):
|
||||
await manager.init_temporary_runtime_llm_model(TEST_EXECUTION_CONTEXT, info('api'))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_temporary_model_rejects_conflicting_identities(manager):
|
||||
with pytest.raises(ValueError):
|
||||
await manager.init_temporary_runtime_llm_model(TEST_EXECUTION_CONTEXT, info('saved', uuid='foreign'))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_key_temporary_model_preserves_inline_configuration(manager):
|
||||
model = await manager.init_temporary_runtime_llm_model(
|
||||
TEST_EXECUTION_CONTEXT,
|
||||
{
|
||||
'name': 'api-model',
|
||||
'provider': {
|
||||
'requester': 'fake-requester',
|
||||
'api_keys': ['synthetic-key'],
|
||||
'base_url': 'https://api.example.invalid',
|
||||
},
|
||||
},
|
||||
)
|
||||
assert model.provider.provider_entity.api_keys == ['synthetic-key']
|
||||
assert model.provider.provider_entity.base_url == 'https://api.example.invalid'
|
||||
Reference in New Issue
Block a user