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:
Hyu
2026-09-06 23:31:02 +08:00
committed by GitHub
parent ec63978ecf
commit 0f216a0d4d
52 changed files with 5490 additions and 360 deletions
+106
View File
@@ -0,0 +1,106 @@
"""Exercise Codex provider wiring through a real LangBot process.
The default run does not contact OpenAI. Set LANGBOT_TEST_CODEX_DEVICE_AUTH=1
to also exercise live device start/pending/cancel, without account sign-in.
OAuth exchange and inference behavior are covered by deterministic tests.
"""
from __future__ import annotations
import os
import time
import pytest
pytestmark = pytest.mark.e2e
def test_codex_provider_disconnected_journey(e2e_client):
credentials = {'user': 'codex-e2e@example.com', 'password': 'codex-local-test-password'}
initialized = e2e_client.post('/api/v1/user/init', json=credentials)
assert initialized.status_code == 200, initialized.text
authenticated = e2e_client.post('/api/v1/user/auth', json=credentials)
assert authenticated.status_code == 200, authenticated.text
headers = {'Authorization': f'Bearer {authenticated.json()["data"]["token"]}'}
bootstrap = e2e_client.get('/api/v1/workspaces/bootstrap', headers=headers)
assert bootstrap.status_code == 200, bootstrap.text
headers['X-Workspace-Id'] = bootstrap.json()['data']['workspaces'][0]['workspace']['uuid']
requesters = e2e_client.get('/api/v1/provider/requesters?type=llm', headers=headers)
assert requesters.status_code == 200, requesters.text
codex = next(item for item in requesters.json()['data']['requesters'] if item['name'] == 'openai-codex')
assert codex['spec']['support_type'] == ['llm']
icon = e2e_client.get('/api/v1/provider/requesters/openai-codex/icon')
assert icon.status_code == 200
assert 'image/' in icon.headers['content-type']
base = '/api/v1/provider/providers'
created = e2e_client.post(
base,
headers=headers,
json={'name': 'Codex E2E', 'requester': 'openai-codex', 'base_url': '', 'api_keys': []},
)
assert created.status_code == 200, created.text
provider_path = f'{base}/{created.json()["data"]["uuid"]}'
try:
provider = e2e_client.get(provider_path, headers=headers)
assert provider.status_code == 200, provider.text
data = provider.json()['data']['provider']
assert data['requester'] == 'openai-codex'
assert data['api_keys'] == []
assert data['base_url'] == 'https://chatgpt.com/backend-api/codex'
assert not {'access_token', 'refresh_token', 'id_token'} & data.keys()
status = e2e_client.get(f'{provider_path}/codex/status', headers=headers)
assert status.status_code == 200, status.text
assert status.json()['data']['connected'] is False
assert status.json()['data']['status'] == 'disconnected'
anonymous = e2e_client.post(f'{provider_path}/codex/device', json={})
assert anonymous.status_code == 401
invalid = e2e_client.put(provider_path, headers=headers, json={'base_url': 'https://example.com'})
assert invalid.status_code == 400, invalid.text
invalid_key = e2e_client.put(provider_path, headers=headers, json={'api_keys': ['not-a-codex-key']})
assert invalid_key.status_code == 400, invalid_key.text
scanned = e2e_client.get(f'{provider_path}/scan-models?type=llm', headers=headers)
assert scanned.status_code == 400, scanned.text
assert 'sign in' in scanned.json()['msg'].lower()
renamed = e2e_client.put(provider_path, headers=headers, json={'name': 'Codex renamed'})
assert renamed.status_code == 200, renamed.text
reread = e2e_client.get(provider_path, headers=headers)
assert reread.json()['data']['provider']['name'] == 'Codex renamed'
disconnected = e2e_client.delete(f'{provider_path}/codex/auth', headers=headers)
assert disconnected.status_code == 200, disconnected.text
# Opt-in smoke contacts real OpenAI device endpoints, but never completes
# account sign-in or prints the one-time code/device credentials.
if os.environ.get('LANGBOT_TEST_CODEX_DEVICE_AUTH') == '1':
started = e2e_client.post(f'{provider_path}/codex/device', headers=headers, json={})
assert started.status_code == 200, started.json().get('msg', 'Device start failed')
attempt = started.json()['data']
assert attempt['verification_uri'] == 'https://auth.openai.com/codex/device'
assert isinstance(attempt['user_code'], str) and attempt['user_code']
assert 0 < attempt['expires_at'] - time.time() <= 900
assert not {'access_token', 'refresh_token', 'device_auth_id'} & attempt.keys()
time.sleep(attempt['interval'])
pending = e2e_client.post(
f'{provider_path}/codex/device/poll',
headers=headers,
json={'authorization_id': attempt['authorization_id']},
)
assert pending.status_code == 200
assert pending.json()['data']['status'] == 'pending'
canceled = e2e_client.delete(f'{provider_path}/codex/device/{attempt["authorization_id"]}', headers=headers)
assert canceled.status_code == 200
expired = e2e_client.post(
f'{provider_path}/codex/device/poll',
headers=headers,
json={'authorization_id': attempt['authorization_id']},
)
assert expired.json()['data']['status'] == 'expired'
finally:
deleted = e2e_client.delete(provider_path, headers=headers)
assert deleted.status_code == 200, deleted.text
assert e2e_client.get(provider_path, headers=headers).status_code == 404
+1 -1
View File
@@ -69,7 +69,7 @@ class LangBotProcess:
# Use coverage.py to collect coverage data
# Set COVERAGE_PROCESS_START to enable coverage in subprocess
self._coverage_file = self.work_dir / '.coverage.e2e'
env['COVERAGE_PROCESS_START'] = str(self.project_root / '.coveragerc')
env['COVERAGE_PROCESS_START'] = str(self.work_dir / '.coveragerc')
env['COVERAGE_FILE'] = str(self._coverage_file)
# Create .coveragerc for subprocess
@@ -0,0 +1,445 @@
"""Deterministic OAuth tests using real SQLite CAS writes, never live credentials."""
import asyncio
import base64
import json
import time
from types import SimpleNamespace
from unittest.mock import AsyncMock
import httpx
import pytest
import pytest_asyncio
import sqlalchemy as sa
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.authz import Permission
from langbot.pkg.api.http.context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
from langbot.pkg.entity.persistence.model import CodexCredential
from langbot.pkg.persistence.alembic_runner import run_alembic_stamp, run_alembic_upgrade
from langbot.pkg.provider.modelmgr.codex_auth import CodexAuth, _tokens, validate_config
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
def context(workspace='w', user='u', principal=PrincipalType.ACCOUNT, permitted=True):
return RequestContext(
'i',
0,
'r',
'user_token',
PrincipalContext(principal, account_uuid=user),
WorkspaceContext(
workspace, 'm', 'owner', frozenset({Permission.PROVIDER_SECRET_MANAGE} if permitted else set())
),
)
def jwt(**claims):
return 'test.' + base64.urlsafe_b64encode(json.dumps(claims).encode()).decode().rstrip('=') + '.test'
def token_response(**extra):
return {
'access_token': jwt(**{'https://api.openai.com/auth': {'chatgpt_account_id': 'account'}}),
'refresh_token': 'refresh-secret',
'expires_in': 3600,
**extra,
}
@pytest_asyncio.fixture
async def auth(tmp_path):
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "codex.db"}')
@sa.event.listens_for(engine.sync_engine, 'connect')
def foreign_keys(connection, _):
connection.execute('PRAGMA foreign_keys=ON')
async with engine.begin() as conn:
await conn.execute(
sa.text(
'CREATE TABLE model_providers (uuid VARCHAR(255) PRIMARY KEY, workspace_uuid VARCHAR(36) NOT NULL, requester TEXT, UNIQUE(workspace_uuid, uuid))'
)
)
await conn.execute(
sa.text(
"INSERT INTO model_providers VALUES ('p','w','openai-codex'), ('other','other','openai-codex'), ('api','w','openai-chat-completions')"
)
)
await run_alembic_stamp(engine, '0021_merge_reasoning_config')
await run_alembic_upgrade(engine, '0022_codex_credentials')
async def execute(statement):
async with engine.begin() as conn:
return await conn.execute(statement)
service = CodexAuth(SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=execute)))
service.engine = engine
await execute(
sa.insert(CodexCredential).values(provider_uuid='p', workspace_uuid='w', payload={}, version=0, lease_until=0)
)
try:
yield service
finally:
await engine.dispose()
async def seed(auth, payload):
await auth.ap.persistence_mgr.execute_async(sa.update(CodexCredential).values(payload=payload))
@pytest.mark.asyncio
async def test_migration_upgrade_repeat_fk_cascade(auth):
await run_alembic_upgrade(auth.engine, '0022_codex_credentials')
await run_alembic_stamp(auth.engine, '0021_merge_reasoning_config')
await run_alembic_upgrade(auth.engine, '0022_codex_credentials')
with pytest.raises(sa.exc.IntegrityError):
await auth.ap.persistence_mgr.execute_async(
sa.insert(CodexCredential).values(
provider_uuid='other', workspace_uuid='w', payload={}, version=0, lease_until=0
)
)
await auth.ap.persistence_mgr.execute_async(sa.text("DELETE FROM model_providers WHERE uuid='p'"))
assert await auth._read('w', 'p') is None
from langbot.pkg.persistence.alembic_runner import run_alembic_downgrade
await run_alembic_downgrade(auth.engine, '0021_merge_reasoning_config')
async with auth.engine.connect() as conn:
assert 'codex_credentials' not in await conn.run_sync(lambda sync: sa.inspect(sync).get_table_names())
await run_alembic_upgrade(auth.engine, '0022_codex_credentials')
async with auth.engine.connect() as conn:
assert 'codex_credentials' in await conn.run_sync(lambda sync: sa.inspect(sync).get_table_names())
@pytest.mark.asyncio
async def test_device_pacing_exchange_secrecy_and_user_binding(auth):
auth._post = AsyncMock(
side_effect=[
httpx.Response(200, json={'device_auth_id': 'device-secret', 'usercode': 'CODE', 'interval': '5'}),
httpx.Response(200, json={'authorization_code': 'code-secret', 'code_verifier': 'verifier-secret'}),
httpx.Response(200, json=token_response()),
]
)
start = await auth.start(context(), 'p')
assert set(start) == {'authorization_id', 'user_code', 'interval', 'expires_at', 'verification_uri'}
assert 'device-secret' not in json.dumps(start)
attempt = start['authorization_id']
with pytest.raises(WorkspaceNotFoundError):
await auth.poll(context(user='attacker'), 'p', attempt)
assert (await auth.poll(context(), 'p', attempt))['status'] == 'pending'
assert auth._post.await_count == 1
row = await auth._read('w', 'p')
row['payload']['pending']['next_poll_at'] = 0
await seed(auth, row['payload'])
assert await auth.poll(context(), 'p', attempt) == {'status': 'connected'}
assert await auth.poll(context(), 'p', attempt) == {'status': 'connected'}
exchange = auth._post.call_args.kwargs['data']
assert exchange['grant_type'] == 'authorization_code'
assert exchange['redirect_uri'] == 'https://auth.openai.com/deviceauth/callback'
assert exchange['code_verifier'] == 'verifier-secret'
status = await auth.status(context(), 'p')
assert set(status) == {'status', 'connected', 'expires_at'}
assert 'secret' not in json.dumps(status)
await auth.disconnect(context(), 'p')
assert (await auth._read('w', 'p'))['payload'] == {}
@pytest.mark.asyncio
@pytest.mark.parametrize(
'ctx,provider,error',
[
(context('other'), 'p', WorkspaceNotFoundError),
(context(principal=PrincipalType.API_KEY), 'p', ValueError),
(context(permitted=False), 'p', ValueError),
(context(), 'api', ValueError),
],
)
async def test_auth_tenant_principal_permission_guards(auth, ctx, provider, error):
auth._post = AsyncMock()
with pytest.raises(error):
await auth.start(ctx, provider)
auth._post.assert_not_called()
@pytest.mark.asyncio
async def test_refresh_cross_instance_single_flight_and_rotation(auth):
old = _tokens(token_response())
old['expires_at'] = 0
await seed(auth, {'tokens': old})
entered, release = asyncio.Event(), asyncio.Event()
async def refresh(*args, **kwargs):
entered.set()
await release.wait()
return httpx.Response(200, json=token_response(refresh_token='rotated-secret'))
auth._post = AsyncMock(side_effect=refresh)
other = CodexAuth(auth.ap)
other._post = auth._post
first = asyncio.create_task(auth.access('w', 'p'))
await entered.wait()
second = asyncio.create_task(other.access('w', 'p'))
release.set()
a, b = await asyncio.gather(first, second)
assert a == b
assert a['refresh_token'] == 'rotated-secret'
assert auth._post.await_count == 1
assert (await auth._read('w', 'p'))['payload']['tokens'] == a
@pytest.mark.asyncio
@pytest.mark.parametrize(
'status,error,invalid',
[
(400, 'invalid_grant', True),
(401, 'refresh_token_reused', True),
(429, 'limited', False),
(500, 'secret-upstream-body', False),
(403, 'permission_denied', False),
],
)
async def test_refresh_errors_are_safe_and_transient_preserves_tokens(auth, status, error, invalid):
old = _tokens(token_response())
old['expires_at'] = 0
await seed(auth, {'tokens': old})
auth._post = AsyncMock(
return_value=httpx.Response(status, json={'error': error, 'access_token': 'secret-upstream-body'})
)
with pytest.raises(ValueError) as caught:
await auth.access('w', 'p')
assert 'secret' not in str(caught.value)
payload = (await auth._read('w', 'p'))['payload']
assert bool(payload.get('invalid')) == invalid
assert ('tokens' not in payload) if invalid else payload['tokens'] == old
@pytest.mark.asyncio
@pytest.mark.parametrize('cancel', [False, True])
async def test_disconnect_or_cancel_fences_inflight_exchange(auth, cancel):
old = _tokens(token_response())
await seed(
auth,
{
'tokens': old,
'pending': {
'authorization_id': 'attempt',
'account_uuid': 'u',
'expires_at': time.time() + 100,
'next_poll_at': 0,
'interval': 5,
'device_auth_id': 'device',
'user_code': 'code',
},
},
)
entered, release = asyncio.Event(), asyncio.Event()
async def post(path, **kwargs):
if path.endswith('/token') and path != '/oauth/token':
return httpx.Response(200, json={'authorization_code': 'code', 'code_verifier': 'verifier'})
entered.set()
await release.wait()
return httpx.Response(200, json=token_response())
auth._post = post
task = asyncio.create_task(auth.poll(context(), 'p', 'attempt'))
await entered.wait()
if cancel:
await auth.cancel(context(), 'p', 'attempt')
else:
await auth.disconnect(context(), 'p')
release.set()
with pytest.raises(ValueError, match='cancelled or replaced'):
await task
payload = (await auth._read('w', 'p'))['payload']
assert payload == ({'tokens': old} if cancel else {})
@pytest.mark.asyncio
@pytest.mark.parametrize('status,interval', [(403, 5), (404, 5), (429, 10)])
async def test_device_pending_and_backoff(auth, status, interval):
await seed(
auth,
{
'pending': {
'authorization_id': 'attempt',
'account_uuid': 'u',
'expires_at': time.time() + 100,
'next_poll_at': 0,
'interval': 5,
'device_auth_id': 'device',
'user_code': 'code',
}
},
)
auth._post = AsyncMock(return_value=httpx.Response(status))
assert await auth.poll(context(), 'p', 'attempt') == {'status': 'pending', 'interval': interval}
assert await auth.poll(context(), 'p', 'attempt') == {'status': 'pending', 'interval': interval}
assert auth._post.await_count == 1
@pytest.mark.asyncio
async def test_device_replacement_expiry_and_idempotent_cancel(auth):
auth._post = AsyncMock(return_value=httpx.Response(200, json={'device_auth_id': 'device', 'user_code': 'CODE'}))
first = await auth.start(context(), 'p')
second = await auth.start(context(), 'p')
assert first['authorization_id'] != second['authorization_id']
assert await auth.poll(context(), 'p', first['authorization_id']) == {'status': 'expired'}
await auth.cancel(context(), 'p', first['authorization_id'])
payload = (await auth._read('w', 'p'))['payload']
assert payload['pending']['authorization_id'] == second['authorization_id']
payload['pending']['expires_at'] = 0
await seed(auth, payload)
assert await auth.poll(context(), 'p', second['authorization_id']) == {'status': 'expired'}
assert 'pending' not in (await auth._read('w', 'p'))['payload']
@pytest.mark.asyncio
async def test_device_accepts_issuer_iso_expiry(auth):
from datetime import datetime, timezone
expires = datetime.fromtimestamp(time.time() + 600, timezone.utc).isoformat().replace('+00:00', 'Z')
auth._post = AsyncMock(
return_value=httpx.Response(200, json={'device_auth_id': 'device', 'user_code': 'CODE', 'expires_at': expires})
)
result = await auth.start(context(), 'p')
assert time.time() < result['expires_at'] < time.time() + 900
@pytest.mark.asyncio
async def test_device_expires_in_fallback(auth):
auth._post = AsyncMock(
return_value=httpx.Response(200, json={'device_auth_id': 'device', 'user_code': 'CODE', 'expires_in': 60})
)
result = await auth.start(context(), 'p')
assert time.time() < result['expires_at'] <= time.time() + 60
@pytest.mark.asyncio
@pytest.mark.parametrize('cancel_reads_before_refresh', [False, True])
async def test_cancel_pending_relogin_waits_for_existing_refresh(auth, cancel_reads_before_refresh):
old = _tokens(token_response())
old['expires_at'] = 0
await seed(auth, {'tokens': old, 'pending': {'authorization_id': 'attempt', 'account_uuid': 'u'}})
entered, release, cancel_read = asyncio.Event(), asyncio.Event(), asyncio.Event()
async def refresh(*args, **kwargs):
entered.set()
await release.wait()
return httpx.Response(200, json=token_response(refresh_token='rotated-secret'))
other = CodexAuth(auth.ap)
original_read = other._read
async def read(workspace, provider):
row = await original_read(workspace, provider)
cancel_read.set()
if cancel_reads_before_refresh:
await entered.wait()
return row
other._read = read
auth._post = refresh
if cancel_reads_before_refresh:
cancelling = asyncio.create_task(other.cancel(context(), 'p', 'attempt'))
await cancel_read.wait()
refreshing = asyncio.create_task(auth.access('w', 'p'))
await entered.wait()
if not cancel_reads_before_refresh:
cancelling = asyncio.create_task(other.cancel(context(), 'p', 'attempt'))
await cancel_read.wait()
await asyncio.sleep(0.05)
try:
assert not cancelling.done(), 'Cancellation must not revoke the refresh lease'
finally:
release.set()
results = await asyncio.gather(refreshing, cancelling, return_exceptions=True)
assert not any(isinstance(result, Exception) for result in results)
payload = (await auth._read('w', 'p'))['payload']
assert payload['tokens']['refresh_token'] == 'rotated-secret'
assert 'pending' not in payload
@pytest.mark.asyncio
@pytest.mark.parametrize('operation', ['save', 'cancel', 'disconnect', 'acquire', 'release', 'read'])
async def test_credential_database_errors_never_expose_secrets(auth, operation):
import traceback
markers = ['ACCESS-MARKER', 'REFRESH-MARKER', 'DEVICE-MARKER', 'VERIFIER-MARKER']
payload = {
'tokens': {'access_token': markers[0], 'refresh_token': markers[1]},
'pending': {
'authorization_id': 'attempt',
'account_uuid': 'u',
'device_auth_id': markers[2],
'code_verifier': markers[3],
},
}
await seed(auth, payload)
if operation == 'read':
auth.ap.persistence_mgr.execute_async = AsyncMock(
side_effect=sa.exc.StatementError(
'failure', 'SELECT credentials', {'payload': payload}, RuntimeError(markers[0])
)
)
else:
column = 'payload' if operation in ('save', 'cancel', 'disconnect') else 'lease_owner'
condition = ' WHEN NEW.lease_owner IS NULL' if operation == 'release' else ''
# Trigger errors can themselves contain secrets, even for parameter-free writes.
await auth.ap.persistence_mgr.execute_async(
sa.text(
f'CREATE TRIGGER reject_write BEFORE UPDATE OF {column} ON codex_credentials{condition} '
f"BEGIN SELECT RAISE(ABORT, '{' '.join(markers)}'); END"
)
)
with pytest.raises(ValueError, match='credential storage') as caught:
if operation == 'save':
async with auth._lease('w', 'p') as owner:
await auth._save('w', 'p', owner, payload)
elif operation == 'cancel':
await auth.cancel(context(), 'p', 'attempt')
elif operation == 'disconnect':
await auth.disconnect(context(), 'p')
elif operation == 'read':
await auth._read('w', 'p')
else:
async with auth._lease('w', 'p'):
pass
rendered = ''.join(traceback.format_exception(caught.value))
assert all(marker not in rendered for marker in markers)
assert caught.value.__suppress_context__
@pytest.mark.asyncio
async def test_credential_serialization_failure_is_sanitized(auth):
import traceback
class Secret:
def __repr__(self):
return 'SERIALIZATION-SECRET'
with pytest.raises(ValueError, match='credential storage') as caught:
async with auth._lease('w', 'p') as owner:
await auth._save('w', 'p', owner, {'tokens': {'refresh_token': Secret()}})
assert 'SERIALIZATION-SECRET' not in ''.join(traceback.format_exception(caught.value))
assert caught.value.__suppress_context__
def test_token_refresh_fallback_and_config_validation():
old = _tokens(token_response())
refreshed = _tokens({'access_token': 'opaque-access', 'expires_in': 3600}, old)
assert refreshed['refresh_token'] == old['refresh_token']
assert refreshed['connection_id'] == old['connection_id']
for expiry in [float('nan'), float('inf'), -1, 'bad']:
with pytest.raises(ValueError):
_tokens(token_response(expires_in=expiry))
data = {'requester': 'openai-codex'}
validate_config(data)
assert data['api_keys'] == []
for update in [{'base_url': 'https://evil.invalid'}, {'api_keys': ['secret']}]:
with pytest.raises(ValueError):
validate_config({**data, **update})
ordinary = {'requester': 'openai-chat-completions', 'api_keys': ['key'], 'base_url': 'https://custom.invalid'}
before = dict(ordinary)
validate_config(ordinary)
assert ordinary == before
@@ -108,7 +108,7 @@ class TestSQLiteMigrationUpgrade:
await run_alembic_upgrade(sqlite_engine, 'head')
assert await get_alembic_current(sqlite_engine) == _get_script_head()
assert _get_script_head() == '0021_merge_reasoning_config'
assert _get_script_head() == '0022_codex_credentials'
@pytest.mark.asyncio
async def test_upgrade_from_reasoning_config_head_to_merged_head(self, sqlite_engine):
@@ -119,7 +119,7 @@ class TestSQLiteMigrationUpgrade:
await run_alembic_stamp(sqlite_engine, '0018_llm_reasoning_config')
await run_alembic_upgrade(sqlite_engine, 'head')
assert await get_alembic_current(sqlite_engine) == '0021_merge_reasoning_config'
assert await get_alembic_current(sqlite_engine) == '0022_codex_credentials'
@pytest.mark.asyncio
async def test_upgrade_from_baseline_to_head(self, sqlite_engine):
@@ -549,10 +549,12 @@ class TestPostgreSQLWorkspaceMigration:
await conn.run_sync(lambda sync_conn: sa.inspect(sync_conn).get_table_names())
)
assert 'workspaces' not in tables_before_migration
assert 'codex_credentials' not in tables_before_migration
await manager._initialize_managed_schema()
async with postgres_engine.connect() as conn:
assert 'codex_credentials' in await conn.run_sync(lambda sync: sa.inspect(sync).get_table_names())
account = (await conn.execute(text('SELECT uuid, status, source FROM users'))).mappings().one()
workspace = (
(await conn.execute(text('SELECT * FROM workspaces WHERE source = :source'), {'source': 'local'}))
@@ -5,6 +5,7 @@ import logging
import os
import pathlib
import sqlite3
from contextlib import closing
import pytest
import sqlalchemy as sa
@@ -34,7 +35,7 @@ def _manifest_payloads(backup_directory) -> list[dict]:
def _assert_verified_backup(payload: dict) -> None:
backup_path = pathlib.Path(payload['backup_path'])
with sqlite3.connect(f'{backup_path.as_uri()}?mode=ro', uri=True) as connection:
with closing(sqlite3.connect(f'{backup_path.as_uri()}?mode=ro', uri=True)) as connection:
assert connection.execute('PRAGMA quick_check').fetchall() == [('ok',)]
assert connection.execute('SELECT version_num FROM alembic_version').fetchone()[0] == payload['source_revision']
@@ -403,10 +403,12 @@ async def test_persistence_startup_defers_workspace_tables_until_account_upgrade
await conn.run_sync(lambda sync_conn: sa.inspect(sync_conn).get_table_names())
)
assert 'workspaces' not in tables_before_migration
assert 'codex_credentials' not in tables_before_migration
await manager._run_alembic_migrations()
async with engine.connect() as conn:
assert 'codex_credentials' in await conn.run_sync(lambda sync: sa.inspect(sync).get_table_names())
workspace = (
(await conn.execute(sa.text("SELECT * FROM workspaces WHERE source = 'local'"))).mappings().one()
)
@@ -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()
+200
View File
@@ -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'