Merge master into dev/4.11.x and preserve plugin runner architecture

Reconcile migration branches without rewriting published revisions; retain additive Codex, monitoring, provider and platform fixes. Keep dynamic runner schemas and Host ownership, restore compatibility regressions, and preserve safe model-test error handling.
This commit is contained in:
RockChinQ
2026-09-16 08:13:07 +00:00
228 changed files with 22138 additions and 5057 deletions
@@ -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):
@@ -12,6 +12,7 @@ import pytest
from unittest.mock import AsyncMock, MagicMock, Mock, patch
from types import SimpleNamespace
import json
import sqlalchemy
import uuid
from langbot.pkg.api.http.service.bot import BotService
@@ -456,6 +457,87 @@ class TestBotServiceCreateBot:
assert 'use_pipeline_name' not in insert_values
assert bot_uuid is not None # Verify UUID was returned
async def test_failed_apply_keeps_bot_saved_visible_and_retryable(self, tmp_path):
"""A saved UUID remains editable after create/update runtime failures."""
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.service.bot_errors import BotApplyError
from langbot.pkg.entity.persistence.user import User
from langbot.pkg.entity.persistence.workspace import Workspace
from langbot.pkg.persistence.mgr import PersistenceManager
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "bots.db"}')
runtime_bot = SimpleNamespace(enable=True, run=AsyncMock())
ap = SimpleNamespace(
instance_config=SimpleNamespace(data={'system': {'limitation': {'max_bots': -1}}}),
platform_mgr=SimpleNamespace(
load_bot=AsyncMock(
side_effect=[
RuntimeError('Invalid token: original-secret'),
RuntimeError('Invalid token: corrected-secret'),
runtime_bot,
]
),
remove_bot=AsyncMock(),
),
sess_mgr=SimpleNamespace(session_list=[]),
)
ap.persistence_mgr = PersistenceManager(ap)
ap.persistence_mgr.db = SimpleNamespace(get_engine=lambda: engine)
service = BotService(ap)
try:
async with engine.begin() as connection:
await connection.execute(sqlalchemy.text('PRAGMA foreign_keys=ON'))
await connection.run_sync(User.__table__.create)
await connection.run_sync(Workspace.__table__.create)
await connection.run_sync(Bot.__table__.create)
await connection.execute(
sqlalchemy.insert(Workspace).values(
uuid=WORKSPACE_UUID, instance_uuid='instance-a', name='Test', slug='test'
)
)
with pytest.raises(BotApplyError) as create_error:
await service.create_bot(
WORKSPACE_UUID,
{
'name': 'Saved bot',
'description': 'Editable after an adapter failure',
'adapter': 'telegram',
'adapter_config': {'token': 'original-secret'},
'enable': True,
},
)
bot_uuid = create_error.value.bot_uuid
assert str(uuid.UUID(bot_uuid)) == bot_uuid
assert 'original-secret' not in str(create_error.value)
assert 'Invalid token' in str(create_error.value)
saved = await service.get_bot(WORKSPACE_UUID, bot_uuid, include_secret=True)
assert saved['uuid'] == bot_uuid
assert saved['adapter_config'] == {'token': 'original-secret'}
assert await service.get_bot('workspace-b', bot_uuid) is None
assert [bot['uuid'] for bot in await service.get_bots(WORKSPACE_UUID)] == [bot_uuid]
with pytest.raises(BotApplyError) as update_error:
await service.update_bot(WORKSPACE_UUID, bot_uuid, {'adapter_config': {'token': 'corrected-secret'}})
assert update_error.value.bot_uuid == bot_uuid
assert 'corrected-secret' not in str(update_error.value)
saved = await service.get_bot(WORKSPACE_UUID, bot_uuid, include_secret=True)
assert saved['adapter_config'] == {'token': 'corrected-secret'}
await service.update_bot(WORKSPACE_UUID, bot_uuid, {'adapter_config': {'token': 'working-token'}})
saved = await service.get_bot(WORKSPACE_UUID, bot_uuid, include_secret=True)
assert saved['adapter_config'] == {'token': 'working-token'}
assert [bot['uuid'] for bot in await service.get_bots(WORKSPACE_UUID)] == [bot_uuid]
async with engine.connect() as connection:
assert await connection.scalar(sqlalchemy.select(sqlalchemy.func.count()).select_from(Bot)) == 1
assert ap.platform_mgr.load_bot.await_count == 3
assert {call.args[1]['uuid'] for call in ap.platform_mgr.load_bot.await_args_list} == {bot_uuid}
runtime_bot.run.assert_awaited_once()
finally:
await engine.dispose()
class TestBotServiceUpdateBot:
"""Tests for update_bot method."""
@@ -1009,6 +1009,37 @@ class TestMCPServiceTestMCPServer:
# Verify - returns task ID
assert task_id == 123
@pytest.mark.parametrize('refresh_first', [False, True])
async def test_persisted_test_preserves_failure_details(self, refresh_first):
from langbot.pkg.provider.tools.loaders.mcp import MCPSessionStatus
runtime_info = {'status': 'error', 'error_message': 'HTTP 403: access denied'}
session = SimpleNamespace(
status=MCPSessionStatus.CONNECTED if refresh_first else MCPSessionStatus.ERROR,
session=object(),
refresh=AsyncMock(side_effect=RuntimeError('refresh failed')),
start=AsyncMock(side_effect=RuntimeError('Connection failed, please check URL')),
get_runtime_info_dict=Mock(return_value=runtime_info),
)
captured = {}
def create_user_task(coroutine, **kwargs):
captured.update(coroutine=coroutine, context=kwargs['context'])
return SimpleNamespace(id=123)
ap = SimpleNamespace(
tool_mgr=SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=session))),
task_mgr=SimpleNamespace(create_user_task=Mock(side_effect=create_user_task)),
)
service = _service(ap)
service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'existing-server'}))
await service.test_mcp_server(_CONTEXT, 'existing-server', {})
with pytest.raises(RuntimeError, match='Connection failed'):
await captured['coroutine']
assert captured['context'].metadata['runtime_info'] == runtime_info
session.start.assert_awaited_once()
assert session.refresh.await_count == int(refresh_first)
async def test_test_mcp_server_not_found_raises(self):
"""Raises ValueError when server not found."""
# Setup
@@ -1052,6 +1083,45 @@ class TestMCPServiceTestMCPServer:
ap.tool_mgr.mcp_tool_loader.load_mcp_server.assert_called_once()
assert task_id == 456
async def test_transient_test_preserves_runtime_info_after_connection_failure(self):
runtime_info = {
'status': 'error',
'error_phase': 'oauth_required',
'retry_count': 1,
}
mock_session = SimpleNamespace(
server_name='oauth-server',
start=AsyncMock(side_effect=RuntimeError('connection failed')),
get_runtime_info_dict=Mock(return_value=runtime_info),
shutdown=AsyncMock(),
)
ap = SimpleNamespace(
tool_mgr=SimpleNamespace(
mcp_tool_loader=SimpleNamespace(load_mcp_server=AsyncMock(return_value=mock_session))
)
)
captured: dict = {}
def create_user_task(coroutine, **kwargs):
captured['coroutine'] = coroutine
captured['context'] = kwargs['context']
return SimpleNamespace(id=457)
ap.task_mgr = SimpleNamespace(create_user_task=Mock(side_effect=create_user_task))
service = _service(ap)
task_id = await service.test_mcp_server(
_CONTEXT,
'_',
{'name': 'OAuth server', 'mode': 'remote', 'enable': True, 'extra_args': {}},
)
assert task_id == 457
with pytest.raises(RuntimeError, match='connection failed'):
await captured['coroutine']
assert captured['context'].metadata['runtime_info'] == runtime_info
mock_session.shutdown.assert_awaited_once_with()
async def test_rejected_transient_test_session_is_shut_down(self):
ap = SimpleNamespace()
mock_session = MagicMock()
@@ -0,0 +1,19 @@
"""Identifier normalization must not rely on SQLite's permissive codecs."""
import pytest
from langbot.pkg.api.http.service import monitoring
@pytest.mark.parametrize(
('value', 'expected'),
[(None, None), ('', ''), ('00123', '00123'), (' 用户 ', ' 用户 '), (123, '123'), (-123, '-123'), (0, '0')],
)
def test_normalize_user_id_preserves_opaque_strings(value, expected):
assert monitoring._normalize_user_id(value) == expected
@pytest.mark.parametrize('value', [True, False, 1.5, b'123', ['123'], {'id': 123}])
def test_normalize_user_id_rejects_unsupported_types(value):
with pytest.raises(TypeError, match='user_id must be a string, integer, or None'):
monitoring._normalize_user_id(value)
@@ -0,0 +1,220 @@
"""Bot-scoped session regressions exercised against real SQL databases."""
import datetime as dt
import logging
from types import SimpleNamespace
import pytest
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.service.monitoring import MonitoringService
from langbot.pkg.entity.persistence.base import Base
from langbot.pkg.entity.persistence import monitoring as models
from langbot.pkg.persistence.mgr import PersistenceManager
from langbot.pkg.pipeline.monitoring_helper import MonitoringHelper
from tests.integration.persistence.test_monitoring_postgres import cloud_database # noqa: F401
pytestmark = pytest.mark.asyncio
@pytest.mark.asyncio(loop_scope='module')
async def test_postgres_upgrade_rls_and_concurrent_bot_counts(cloud_database): # noqa: F811
import asyncio
import importlib
from alembic.migration import MigrationContext
from alembic.operations import Operations
from tests.integration.persistence.test_monitoring_postgres import WORKSPACE_A, _context, _read
ap, admin = cloud_database
service = ap.monitoring_service
ctx = _context(WORKSPACE_A)
await service.record_session_start(ctx, session_id='person_42', **resource('a'))
for bot in ['a', 'b']:
await service.record_message(ctx, session_id='person_42', message_content=bot, **resource(bot))
async with admin.begin() as conn:
def migrate(connection):
migration = importlib.import_module('langbot.pkg.persistence.alembic.versions.0023_bot_scoped_sessions')
with Operations.context(MigrationContext.configure(connection)):
migration.downgrade()
migration.upgrade()
rls = connection.execute(
sa.text("SELECT relrowsecurity, relforcerowsecurity FROM pg_class WHERE relname='monitoring_sessions'")
).one()
assert tuple(rls) == (True, True)
assert (
connection.execute(
sa.text("SELECT count(*) FROM pg_policies WHERE tablename='monitoring_sessions'")
).scalar_one()
== 1
)
await conn.run_sync(migrate)
rows, total = await _read(service, 'get_sessions', ctx)
assert total == 2
assert {r['bot_id']: r['message_count'] for r in rows} == {'a': 1, 'b': 1}
await asyncio.gather(*[service.record_session_start(ctx, session_id='race', **resource('a')) for _ in range(10)])
result = await _read(service, 'get_session_analysis', ctx, 'race', bot_id='a')
assert result['session']['message_count'] == 10
assert not (await _read(service, 'get_session_analysis', ctx, 'person_42'))['found']
assert (await _read(service, 'get_session_analysis', ctx, 'person_42', bot_id='b'))['message_stats']['total'] == 1
async def test_migration_reconstructs_collisions_and_preserves_indexes(service):
import importlib
from alembic.migration import MigrationContext
from alembic.operations import Operations
engine = service.ap.persistence_mgr.get_db_engine()
async with engine.begin() as conn:
def upgrade(connection):
table = models.MonitoringSession.__table__
table.drop(connection)
metadata = sa.MetaData()
legacy = table.to_metadata(metadata)
legacy.primary_key._columns.remove(legacy.c.bot_id)
legacy.c.bot_id.primary_key = False
# Resolve the unchanged Workspace FK in copied metadata.
Base.metadata.tables['workspaces'].to_metadata(metadata)
legacy.create(connection)
now = dt.datetime(2026, 1, 1)
connection.execute(
sa.insert(legacy).values(
workspace_uuid='workspace',
session_id='person_42',
**resource('a'),
message_count=99,
start_time=now,
last_activity=now,
is_active=True,
)
)
for bot in ['a', 'b']:
connection.execute(
sa.insert(models.MonitoringMessage).values(
id=bot,
workspace_uuid='workspace',
timestamp=now,
**resource(bot),
session_id='person_42',
message_content=bot,
role='user',
status='success',
level='info',
)
)
indexes = {i['name'] for i in sa.inspect(connection).get_indexes('monitoring_sessions')}
migration = importlib.import_module('langbot.pkg.persistence.alembic.versions.0023_bot_scoped_sessions')
with Operations.context(MigrationContext.configure(connection)):
migration.upgrade()
migration.upgrade() # Fresh/already-upgraded schema is safe.
assert sa.inspect(connection).get_pk_constraint('monitoring_sessions')['constrained_columns'] == [
'workspace_uuid',
'bot_id',
'session_id',
]
assert indexes <= {i['name'] for i in sa.inspect(connection).get_indexes('monitoring_sessions')}
await conn.run_sync(upgrade)
rows, total = await service.get_sessions(context())
assert total == 2
assert {r['bot_id']: r['message_count'] for r in rows} == {'a': 1, 'b': 1}
assert {r['pipeline_id'] for r in rows} == {'a', 'b'}
def context(bot=None):
return ExecutionContext(instance_uuid='test', workspace_uuid='workspace', placement_generation=1, bot_uuid=bot)
def resource(bot):
return dict(bot_id=bot, bot_name=bot, pipeline_id=bot, pipeline_name=bot)
@pytest.fixture
async def service():
engine = create_async_engine('sqlite+aiosqlite:///:memory:')
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
class Persistence:
serialize_model = PersistenceManager.serialize_model
def get_db_engine(self):
return engine
async def execute_async(self, stmt):
async with engine.begin() as conn:
return await conn.execute(stmt)
ap = SimpleNamespace(persistence_mgr=Persistence(), logger=logging.getLogger(__name__))
ap.monitoring_service = MonitoringService(ap)
yield ap.monitoring_service
await engine.dispose()
async def test_helper_first_message_count_and_two_bot_isolation(service):
for bot in ['a', 'b', 'a']:
query = SimpleNamespace(
_execution_context=context(bot),
launcher_type='person',
launcher_id=42,
sender_id=42,
message_chain=SimpleNamespace(model_dump=lambda: []),
)
assert await MonitoringHelper.record_query_start(service.ap, query, **resource(bot))
rows, total = await service.get_sessions(context())
assert total == 2
assert {r['bot_id']: r['message_count'] for r in rows} == {'a': 2, 'b': 1}
assert {r['pipeline_id'] for r in rows} == {'a', 'b'}
assert {r['session_id'] for r in rows} == {'person_42'}
async def test_analysis_fails_closed_and_scopes_statistics(service):
for bot in ['a', 'b']:
await service.record_session_start(context(bot), session_id='person_42', **resource(bot))
await service.record_message(context(bot), session_id='person_42', message_content=bot, **resource(bot))
assert (await service.get_session_analysis(context(), 'person_42'))['found'] is False
result = await service.get_session_analysis(context(), 'person_42', bot_id='b')
assert result['message_stats']['total'] == 1
assert result['session']['bot_id'] == 'b'
async def test_activity_requires_bot_and_upsert_counts_racing_first_queries(service):
for _ in range(2):
await service.record_session_start(context('a'), session_id='person_42', **resource('a'))
with pytest.raises(ValueError, match='bot'):
await service.update_session_activity(context(), 'person_42')
assert await service.update_session_activity(context('a'), 'person_42')
assert not await service.update_session_activity(context('b'), 'person_42')
rows, _ = await service.get_sessions(context())
assert rows[0]['message_count'] == 3
async def test_old_active_sessions_are_listed_exported_and_not_cleaned(service):
for bot in ['a', 'b']:
await service.record_session_start(context(bot), session_id='person_42', **resource(bot))
old = dt.datetime(2000, 1, 1)
await service.ap.persistence_mgr.execute_async(sa.update(models.MonitoringSession).values(start_time=old))
await service.ap.persistence_mgr.execute_async(
sa.update(models.MonitoringSession).where(models.MonitoringSession.bot_id == 'a').values(last_activity=old)
)
since = dt.datetime.now(dt.timezone.utc).replace(tzinfo=None) - dt.timedelta(days=1)
rows, total = await service.get_sessions(context(), start_time=since)
assert total == 1 and rows[0]['bot_id'] == 'b'
assert len(await service.export_sessions(context(), start_time=since)) == 1
count = await service._delete_expired_in_batches(
context(),
models.MonitoringSession,
models.MonitoringSession.last_activity,
models.MonitoringSession.session_id,
since,
1,
2,
)
assert count == 1
rows, total = await service.get_sessions(context())
assert total == 1 and rows[0]['bot_id'] == 'b'
@@ -138,6 +138,39 @@ async def test_same_session_and_resource_ids_do_not_collide(service):
assert (await service.get_message_details(context_a, message_b))['found'] is False
async def test_session_search_matches_user_id_or_name_within_workspace(service):
context_a = _context(WORKSPACE_A)
context_b = _context(WORKSPACE_B)
fixtures = [
(context_a, 'session-id-match', 'customer-42', 'Alice'),
(context_a, 'session-name-match', 'customer-99', 'Bob Alice Cooper'),
(context_a, 'session-no-match', 'customer-7', 'Bob'),
(context_b, 'session-other-workspace', 'customer-42', 'Alice'),
]
for context, session_id, user_id, user_name in fixtures:
await service.record_session_start(
context,
session_id=session_id,
bot_id='same-bot',
bot_name='Same Bot',
pipeline_id='same-pipeline',
pipeline_name='Same Pipeline',
user_id=user_id,
user_name=user_name,
)
by_id, id_total = await service.get_sessions(context_a, user_query='customer-42')
by_name, name_total = await service.get_sessions(context_a, user_query='alice')
assert id_total == 1
assert [session['session_id'] for session in by_id] == ['session-id-match']
assert name_total == 2
assert {session['session_id'] for session in by_name} == {
'session-id-match',
'session-name-match',
}
async def test_tool_call_inherits_context_from_connection_message_row(service):
context = _context(WORKSPACE_A)
message_id = await _record_message(service, context, 'tool context')
@@ -0,0 +1,125 @@
from __future__ import annotations
import datetime
from types import SimpleNamespace
import pytest
import sqlalchemy
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.entity.persistence.base import Base
from langbot.pkg.entity.persistence.monitoring import MonitoringLLMCall, MonitoringMessage
from langbot.pkg.entity.persistence.workspace import Workspace
pytestmark = pytest.mark.asyncio
A = '00000000-0000-0000-0000-00000000000a'
B = '00000000-0000-0000-0000-00000000000b'
START = datetime.datetime(2026, 1, 1)
@pytest.fixture
async def traffic_app():
engine = create_async_engine('sqlite+aiosqlite:///:memory:')
async with engine.begin() as connection:
await connection.run_sync(Base.metadata.create_all)
await connection.execute(
sqlalchemy.insert(Workspace),
[
{'uuid': wid, 'instance_uuid': 'instance', 'name': wid, 'slug': wid, 'source': 'cloud_projection'}
for wid in (A, B)
],
)
for wid, bot, count in [(A, 'bot-a', 60), (A, 'bot-b', 7), (B, 'bot-a', 9)]:
common = {
'workspace_uuid': wid,
'timestamp': START,
'bot_id': bot,
'bot_name': bot,
'pipeline_id': 'pipeline',
'pipeline_name': 'Pipeline',
'session_id': 'person_42',
'status': 'success',
}
await connection.execute(
sqlalchemy.insert(MonitoringMessage),
[
dict(common, id=f'{wid}-{bot}-{i}', message_content='test fixture', level='info', role='user')
for i in range(count)
],
)
await connection.execute(
sqlalchemy.insert(MonitoringLLMCall),
[
dict(
common,
id=f'{wid}-{bot}-{i}',
model_name='fixture-model',
input_tokens=1,
output_tokens=1,
total_tokens=2,
duration=1,
)
for i in range(count)
],
)
class Persistence:
def get_db_engine(self):
return engine
async def execute_async(self, statement):
async with engine.connect() as connection:
return await connection.execute(statement)
yield SimpleNamespace(persistence_mgr=Persistence())
await engine.dispose()
async def test_traffic_counts_all_rows_not_just_latest_page(traffic_app):
from langbot.pkg.api.http.service.monitoring_traffic import get_traffic_series
context = ExecutionContext(instance_uuid='instance', workspace_uuid=A, placement_generation=1)
result = await get_traffic_series(
traffic_app, context, bot_ids=['bot-a'], start_time=START, end_time=START + datetime.timedelta(hours=2)
)
assert result['bucket'] == 'hour'
assert result['truncated'] is False
assert sum(point['messages'] for point in result['points']) == 60
assert sum(point['llm_calls'] for point in result['points']) == 60
assert len(result['points']) == 3
assert result['points'][1]['messages'] == result['points'][1]['llm_calls'] == 0
assert result['points'][0]['timestamp'] == '2026-01-01T00:00:00Z'
async def test_traffic_workspace_pipeline_and_empty_filters(traffic_app):
from langbot.pkg.api.http.service.monitoring_traffic import get_traffic_series
context = ExecutionContext(instance_uuid='instance', workspace_uuid=B, placement_generation=1)
kwargs = dict(start_time=START, end_time=START + datetime.timedelta(hours=2))
result = await get_traffic_series(traffic_app, context, **kwargs)
assert sum(point['messages'] for point in result['points']) == 9
empty = await get_traffic_series(traffic_app, context, pipeline_ids=['missing'], **kwargs)
assert sum(point['messages'] for point in empty['points']) == 0
assert sum(point['llm_calls'] for point in empty['points']) == 0
async def test_traffic_bounds_large_ranges_and_marks_truncation(traffic_app):
from langbot.pkg.api.http.service.monitoring_traffic import get_traffic_series
context = ExecutionContext(instance_uuid='instance', workspace_uuid=A, placement_generation=1)
result = await get_traffic_series(
traffic_app, context, start_time=START, end_time=START + datetime.timedelta(days=5000)
)
assert result['bucket'] == 'day'
assert result['truncated'] is True
assert len(result['points']) == 1000
async def test_traffic_fails_closed_without_workspace(traffic_app):
from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.service.monitoring_traffic import get_traffic_series
with pytest.raises(WorkspaceRequiredError):
await get_traffic_series(traffic_app, 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')))
@@ -95,7 +95,9 @@ class TestSpaceServiceGetOAuthAuthorizeUrl:
result = service.get_oauth_authorize_url('http://localhost/callback')
# Verify
assert parse_qs(urlsplit(result).query)['redirect_uri'] == ['http://localhost/callback']
query = parse_qs(urlsplit(result).query)
assert query['redirect_uri'] == ['http://localhost/callback']
assert query['code_contract'] == ['redirect-v1']
assert 'https://space.langbot.app/auth/authorize' in result
def test_get_oauth_authorize_url_with_state(self):
@@ -578,12 +580,14 @@ class TestSpaceServiceExchangeOAuthCode:
'auth_code',
['workspace-1'],
{'workspace-1': 1_700_000_000},
redirect_uri='https://oss.example/auth/space/callback',
)
# Verify
assert result['access_token'] == 'new_access_token'
assert mock_session_obj.post.call_args.kwargs['json'] == {
'code': 'auth_code',
'redirect_uri': 'https://oss.example/auth/space/callback',
'instance_id': constants.instance_id,
'workspace_uuids': ['workspace-1'],
'workspace_created_ats': {'workspace-1': 1_700_000_000},
@@ -823,7 +827,9 @@ class TestSpaceServiceGetModels:
class TestSpaceServiceGetModelSelection:
"""Tests for availability-ranked model selection."""
@pytest.mark.parametrize('response_shape', ['direct', 'models-envelope', 'availability-wrapper'])
@pytest.mark.parametrize(
'response_shape', ['direct', 'models-envelope', 'availability-wrapper', 'legacy-flat-wrapper']
)
async def test_preserves_selection_order_and_category_query(self, response_shape):
ap = SimpleNamespace(instance_config=SimpleNamespace(data={}))
service = SpaceService(ap)
@@ -859,6 +865,8 @@ class TestSpaceServiceGetModelSelection:
}
for index, model in enumerate(models)
]
elif response_shape == 'legacy-flat-wrapper':
data = [{'model': model, 'latency_ms': index + 10, 'http_code': 200} for index, model in enumerate(models)]
else:
data = models
payload = {'code': 0, 'data': data}
@@ -0,0 +1,103 @@
"""
Unit tests for Passkey WebAuthn service operations in UserService.
"""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.service.user import UserService
from langbot.pkg.entity.persistence.user import AccountStatus, User
pytestmark = pytest.mark.asyncio
class TestPasskeyChallengeLifecycle:
async def test_challenge_issuance_and_consumption(self):
service = UserService(SimpleNamespace())
token, challenge_bytes = await service.issue_passkey_challenge(
purpose='register',
rp_id='localhost',
origin='http://localhost:3000',
account_uuid='acc-123',
user_email='user@example.com',
)
assert len(token) > 20
assert len(challenge_bytes) == 32
data = await service.consume_passkey_challenge(token, 'register')
assert data.challenge == challenge_bytes
assert data.rp_id == 'localhost'
assert data.origin == 'http://localhost:3000'
assert data.account_uuid == 'acc-123'
assert data.user_email == 'user@example.com'
# Replay should fail
with pytest.raises(ValueError, match='Invalid or expired passkey challenge'):
await service.consume_passkey_challenge(token, 'register')
async def test_challenge_purpose_mismatch_fails(self):
service = UserService(SimpleNamespace())
token, _ = await service.issue_passkey_challenge(
purpose='register',
rp_id='localhost',
origin='http://localhost:3000',
)
with pytest.raises(ValueError, match='Passkey challenge purpose mismatch'):
await service.consume_passkey_challenge(token, 'auth')
async def test_challenge_expiration(self):
service = UserService(SimpleNamespace())
token, _ = await service.issue_passkey_challenge(
purpose='auth',
rp_id='localhost',
origin='http://localhost:3000',
ttl_seconds=0,
)
with pytest.raises(ValueError, match='Invalid or expired passkey challenge'):
await service.consume_passkey_challenge(token, 'auth')
class TestPasskeyOptionsGeneration:
async def test_generate_registration_options(self):
service = UserService(SimpleNamespace())
mock_user = Mock(spec=User)
mock_user.uuid = 'acc-test-uuid'
mock_user.user = 'test@example.com'
mock_user.status = AccountStatus.ACTIVE.value
service.get_user_by_uuid = AsyncMock(return_value=mock_user)
service.get_user_passkeys = AsyncMock(return_value=[])
options, token = await service.generate_passkey_registration_options(
account_uuid='acc-test-uuid',
rp_id='localhost',
origin='http://localhost:3000',
rp_name='LangBot Test',
)
assert isinstance(options, dict)
assert options['rp']['name'] == 'LangBot Test'
assert options['rp']['id'] == 'localhost'
assert options['user']['name'] == 'test@example.com'
assert 'challenge' in options
assert len(token) > 0
async def test_generate_authentication_options_discoverable(self):
service = UserService(SimpleNamespace())
options, token = await service.generate_passkey_authentication_options(
rp_id='localhost',
origin='http://localhost:3000',
)
assert isinstance(options, dict)
assert options['rpId'] == 'localhost'
assert 'challenge' in options
assert len(token) > 0
@@ -0,0 +1,182 @@
from __future__ import annotations
import copy
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
from langbot.pkg.api.http.controller.groups.provider.models import (
EmbeddingModelsRouterGroup,
LLMModelsRouterGroup,
RerankModelsRouterGroup,
)
from langbot.pkg.api.http.controller.groups.provider.providers import ModelProvidersRouterGroup
from langbot.pkg.api.http.controller.groups.provider.query import resolve_include_secret
from langbot.pkg.api.http.service.secrets import redact_secrets
pytestmark = pytest.mark.asyncio
RAW_PROVIDER = {
'uuid': 'provider-test',
'name': 'Test Provider',
'api_keys': ['provider-secret'],
}
RAW_MODEL = {
'uuid': 'model-test',
'name': 'Test Model',
'extra_args': {'headers': {'Authorization': 'Bearer model-secret'}},
}
def _access(role: str):
return SimpleNamespace(
execution=SimpleNamespace(instance_uuid='instance-test', placement_generation=1),
workspace=SimpleNamespace(uuid='workspace-test'),
membership=SimpleNamespace(uuid='membership-test', role=role, projection_revision=1),
)
def _project(value: dict, include_secret: bool) -> dict:
value = copy.deepcopy(value)
return value if include_secret else redact_secrets(value)
async def _create_client(role: str):
application = SimpleNamespace()
account = SimpleNamespace(uuid='account-test', user='test@example.com')
application.user_service = SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account))
application.apikey_service = SimpleNamespace(authenticate_api_key=AsyncMock(return_value=None))
application.workspace_collaboration_service = SimpleNamespace(
resolve_account_workspace=AsyncMock(return_value=_access(role))
)
async def get_providers(_context, *, include_secret=False):
return [_project(RAW_PROVIDER, include_secret)]
async def get_provider(_context, _uuid, *, include_secret=False):
return _project(RAW_PROVIDER, include_secret)
application.provider_service = SimpleNamespace(
get_providers=AsyncMock(side_effect=get_providers),
get_provider=AsyncMock(side_effect=get_provider),
get_provider_model_counts=AsyncMock(
return_value={'llm_count': 1, 'embedding_count': 1, 'rerank_count': 1}
),
)
def model_service(list_name: str, get_name: str):
async def get_models(_context, *, include_secret=False):
return [_project(RAW_MODEL, include_secret)]
async def get_model(_context, _uuid, *, include_secret=False):
return _project(RAW_MODEL, include_secret)
return SimpleNamespace(
**{
list_name: AsyncMock(side_effect=get_models),
get_name: AsyncMock(side_effect=get_model),
}
)
application.llm_model_service = model_service('get_llm_models', 'get_llm_model')
application.embedding_models_service = model_service('get_embedding_models', 'get_embedding_model')
application.rerank_models_service = model_service('get_rerank_models', 'get_rerank_model')
quart_app = quart.Quart(__name__)
for router_type in (
ModelProvidersRouterGroup,
LLMModelsRouterGroup,
EmbeddingModelsRouterGroup,
RerankModelsRouterGroup,
):
await router_type(application, quart_app).initialize()
return application, quart_app.test_client()
def _headers() -> dict[str, str]:
return {'Authorization': 'Bearer test-token'}
@pytest.mark.parametrize(
('raw_value', 'permitted', 'expected', 'error'),
[
(None, True, True, None),
(None, False, False, None),
('false', True, False, None),
('true', True, True, None),
('true', False, False, None),
('invalid', True, False, 'include_secret must be either true or false'),
],
)
def test_resolve_include_secret(raw_value, permitted, expected, error):
assert resolve_include_secret(raw_value, permitted=permitted) == (expected, error)
@pytest.mark.parametrize(
'endpoint',
[
'/api/v1/provider/providers',
'/api/v1/provider/models/llm',
'/api/v1/provider/models/embedding',
'/api/v1/provider/models/rerank',
],
)
async def test_default_preserves_secrets_and_explicit_false_redacts_high_permission_reads(endpoint):
application, client = await _create_client('developer')
default_response = await client.get(endpoint, headers=_headers())
false_response = await client.get(f'{endpoint}?include_secret=false', headers=_headers())
assert default_response.status_code == 200
assert false_response.status_code == 200
default_data = await default_response.get_json()
false_data = await false_response.get_json()
default_value = default_data['data'].get('providers', default_data['data'].get('models'))[0]
false_value = false_data['data'].get('providers', false_data['data'].get('models'))[0]
assert '***' not in str(default_value)
assert '***' in str(false_value)
@pytest.mark.parametrize(
'endpoint',
[
'/api/v1/provider/providers',
'/api/v1/provider/models/llm',
'/api/v1/provider/models/embedding',
'/api/v1/provider/models/rerank',
],
)
async def test_explicit_true_does_not_grant_low_permission_reads(endpoint):
_application, client = await _create_client('viewer')
response = await client.get(f'{endpoint}?include_secret=true', headers=_headers())
assert response.status_code == 200
data = await response.get_json()
value = data['data'].get('providers', data['data'].get('models'))[0]
assert '***' in str(value)
@pytest.mark.parametrize(
'endpoint',
[
'/api/v1/provider/providers',
'/api/v1/provider/providers/provider-test',
'/api/v1/provider/models/llm',
'/api/v1/provider/models/llm/model-test',
'/api/v1/provider/models/embedding',
'/api/v1/provider/models/embedding/model-test',
'/api/v1/provider/models/rerank',
'/api/v1/provider/models/rerank/model-test',
],
)
async def test_invalid_include_secret_returns_bad_request(endpoint):
_application, client = await _create_client('developer')
response = await client.get(f'{endpoint}?include_secret=maybe', headers=_headers())
assert response.status_code == 400
assert (await response.get_json())['msg'] == 'include_secret must be either true or false'
@@ -0,0 +1,302 @@
"""Regression tests for recovery-key hardening (#2392).
Covers two attack surfaces reported in GHSA-4xcp-6758-rxqv:
1. ``genkeys.py`` generated ``system.recovery_key`` with only 24 bits of
entropy (``secrets.token_hex(3)``), making the whole keyspace brute-forceable.
2. ``POST /api/v1/user/reset-password`` (unauthenticated) checked its failure
counter across ``await`` points, so concurrent guesses all passed the gate
before any accounting happened; admission is now a synchronous fixed-window
quota consumed at entry, plus constant-time key comparison.
"""
from __future__ import annotations
import asyncio
import logging
import time
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
from langbot.pkg.api.http.controller.groups import user as user_module
from langbot.pkg.api.http.controller.groups.user import UserRouterGroup
from langbot.pkg.core.stages.genkeys import GenKeysStage
pytestmark = pytest.mark.asyncio
STORED_KEY = 'ABCD2345'
@pytest.fixture(autouse=True)
def _reset_quota_state():
"""Reset the module-level admission-quota state before each test."""
user_module._reset_password_state['window_started_at'] = 0.0
user_module._reset_password_state['attempts'] = 0
yield
user_module._reset_password_state['window_started_at'] = 0.0
user_module._reset_password_state['attempts'] = 0
@pytest.fixture(autouse=True)
def _fast_sleep(monkeypatch):
"""Neutralize the fixed 3s delay so tests run instantly."""
monkeypatch.setattr(user_module, 'asyncio', SimpleNamespace(sleep=AsyncMock()))
# ---------------------------------------------------------------------------
# genkeys.py: recovery-key generation and compatibility
# ---------------------------------------------------------------------------
def _make_genkeys_ap(existing_key: str) -> SimpleNamespace:
"""Build a minimal Application mock for GenKeysStage.
Mirrors the real boot order: no ``logger`` attribute is set because
GenKeysStage runs before SetupLoggerStage.
"""
return SimpleNamespace(
instance_config=SimpleNamespace(
data={'system': {'jwt': {'secret': 'jwt-secret'}, 'recovery_key': existing_key}},
dump_config=AsyncMock(),
),
)
async def test_recovery_key_generation_is_short_and_unambiguous():
"""Eight random base32 characters balance manual entry and online throttling."""
ap = _make_genkeys_ap(existing_key='')
await GenKeysStage().run(ap)
key = ap.instance_config.data['system']['recovery_key']
assert len(key) == 8
assert set(key) <= set('23456789ABCDEFGHJKLMNPQRSTUVWXYZ')
assert ap.instance_config.dump_config.called
async def test_legacy_low_entropy_key_preserved_with_warning(caplog):
"""A legacy 6-char key must keep working but emit a warning, without ap.logger."""
ap = _make_genkeys_ap(existing_key='ABC123')
with caplog.at_level(logging.WARNING, logger='langbot.pkg.core.stages.genkeys'):
await GenKeysStage().run(ap)
assert ap.instance_config.data['system']['recovery_key'] == 'ABC123'
assert any('Low-entropy' in record.message for record in caplog.records)
assert not ap.instance_config.dump_config.called
@pytest.mark.parametrize('existing_key', ['ABC123', 'ABCD2345', 'aB-_' * 10 + 'xYz', '自定义恢复密钥'])
async def test_recovery_key_generation_preserves_existing_key(existing_key):
"""An explicitly configured recovery key must not be regenerated on boot."""
ap = _make_genkeys_ap(existing_key=existing_key)
await GenKeysStage().run(ap)
assert ap.instance_config.data['system']['recovery_key'] == existing_key
assert not ap.instance_config.dump_config.called
async def test_generated_key_is_preserved_without_legacy_warning(caplog):
"""A restart must not warn about or replace the new eight-character key."""
ap = _make_genkeys_ap(existing_key='')
await GenKeysStage().run(ap)
key = ap.instance_config.data['system']['recovery_key']
assert len(key) == 8
ap.instance_config.dump_config.reset_mock()
with caplog.at_level(logging.WARNING, logger='langbot.pkg.core.stages.genkeys'):
await GenKeysStage().run(ap)
assert ap.instance_config.data['system']['recovery_key'] == key
assert not caplog.records
ap.instance_config.dump_config.assert_not_awaited()
async def test_eight_character_key_does_not_trigger_legacy_warning(caplog):
ap = _make_genkeys_ap(existing_key='ABCD2345')
with caplog.at_level(logging.WARNING, logger='langbot.pkg.core.stages.genkeys'):
await GenKeysStage().run(ap)
assert not caplog.records
# ---------------------------------------------------------------------------
# POST /api/v1/user/reset-password: admission quota + constant-time compare
# ---------------------------------------------------------------------------
async def _create_client(stored_key: str = STORED_KEY):
"""Create a Quart test client with a mocked Application."""
quart_app = quart.Quart(__name__)
user_obj = SimpleNamespace(uuid='user-uuid', user='admin@example.com')
reset_password = AsyncMock()
get_user_by_email = AsyncMock(return_value=user_obj)
ap = SimpleNamespace(
user_service=SimpleNamespace(
is_initialized=AsyncMock(return_value=True),
get_user_by_email=get_user_by_email,
reset_password=reset_password,
),
instance_config=SimpleNamespace(
data={'system': {'recovery_key': stored_key}},
),
)
router = UserRouterGroup(ap, quart_app)
await router.initialize()
client = quart_app.test_client()
return client, reset_password, get_user_by_email
def _payload(key: str = STORED_KEY) -> dict:
return {'user': 'admin@example.com', 'recovery_key': key, 'new_password': 'NewPass1!'}
@pytest.mark.parametrize('key', [STORED_KEY, 'ABC123', 'aB-_' * 10 + 'xYz', '自定义恢复密钥'])
async def test_correct_key_resets_password(key):
"""New, legacy and explicitly configured keys all remain usable verbatim."""
client, reset_password, _ = await _create_client(stored_key=key)
resp = await client.post('/api/v1/user/reset-password', json=_payload(key))
assert resp.status_code == 200
assert (await resp.get_json())['code'] == 0
reset_password.assert_awaited_once_with('admin@example.com', 'NewPass1!')
async def test_wrong_key_rejected_without_reset():
"""A wrong recovery key returns 403 and never touches the password."""
client, reset_password, _ = await _create_client()
resp = await client.post('/api/v1/user/reset-password', json=_payload(key='WRONG'))
assert resp.status_code == 403
reset_password.assert_not_awaited()
async def test_non_string_recovery_key_does_not_crash():
"""Malformed recovery-key payloads must be rejected, not raise a 500.
Constant-time comparison via hmac.compare_digest on bytes requires the
input to be a str; other JSON types must fail closed.
"""
client, reset_password, _ = await _create_client()
resp = await client.post(
'/api/v1/user/reset-password',
json={'user': 'admin@example.com', 'recovery_key': 12345, 'new_password': 'NewPass1!'},
)
assert resp.status_code == 403
reset_password.assert_not_awaited()
@pytest.mark.parametrize('key', ['奇数密钥不是ASCII', '\ud800', '\udfff'])
async def test_non_ascii_recovery_key_does_not_crash(key):
"""Non-ASCII keys must compare safely (encode-based constant-time compare)."""
client, _, _ = await _create_client()
resp = await client.post(
'/api/v1/user/reset-password',
json={'user': 'admin@example.com', 'recovery_key': key, 'new_password': 'NewPass1!'},
)
assert resp.status_code == 403
async def test_quota_exhausted_after_max_attempts():
"""After MAX admitted attempts even a correct key must be rejected with 429 (#2392).
Every admission consumes quota regardless of outcome; the legacy endpoint
accepted every guess independently, exhausting the 24-bit keyspace via bursts.
"""
client, reset_password, _ = await _create_client()
for _ in range(user_module._MAX_RESET_ATTEMPTS_PER_WINDOW):
resp = await client.post('/api/v1/user/reset-password', json=_payload(key='WRONG'))
assert resp.status_code == 403
# The very next request carries the CORRECT key but has no quota left.
resp = await client.post('/api/v1/user/reset-password', json=_payload())
assert resp.status_code == 429
reset_password.assert_not_awaited()
async def test_quota_rejects_before_touching_user_lookup():
"""An exhausted quota must reject early, before the sleep and any service calls."""
client, _, get_user_by_email = await _create_client()
user_module._reset_password_state['attempts'] = user_module._MAX_RESET_ATTEMPTS_PER_WINDOW
user_module._reset_password_state['window_started_at'] = time.monotonic()
resp = await client.post('/api/v1/user/reset-password', json=_payload())
assert resp.status_code == 429
get_user_by_email.assert_not_awaited()
async def test_window_rolls_over_and_admits_again():
"""Once the fixed window elapses, the quota resets and a correct key works again."""
client, reset_password, _ = await _create_client()
user_module._reset_password_state['attempts'] = user_module._MAX_RESET_ATTEMPTS_PER_WINDOW
user_module._reset_password_state['window_started_at'] = time.monotonic() - user_module._RESET_WINDOW_SECONDS - 1
resp = await client.post('/api/v1/user/reset-password', json=_payload())
assert resp.status_code == 200
reset_password.assert_awaited_once()
async def test_success_does_not_restore_quota():
"""A successful reset does NOT restore quota: brute-force budget survives wins (#2392).
The legacy clear-on-success let attackers interleave correct-looking states;
success only proves knowledge of the key once, it must not refill attempts.
"""
client, _, _ = await _create_client()
for _ in range(user_module._MAX_RESET_ATTEMPTS_PER_WINDOW - 1):
resp = await client.post('/api/v1/user/reset-password', json=_payload(key='WRONG'))
assert resp.status_code == 403
# Last slot is spent on the genuine reset.
resp = await client.post('/api/v1/user/reset-password', json=_payload())
assert resp.status_code == 200
# Quota is exhausted; even a correct key waits for the next window.
resp = await client.post('/api/v1/user/reset-password', json=_payload())
assert resp.status_code == 429
async def test_concurrent_burst_cannot_bypass_quota(monkeypatch):
"""A 20-request burst yields exactly {403: 5, 429: 15} (#2392 regression).
The vulnerable version accounted failures after several awaits, letting all
concurrent requests pass the gate ({403: 20}). Admission is now synchronous
and await-free, so total admissions are capped regardless of scheduling.
"""
# Swap the AsyncMock sleep for a real cooperative yield so tasks actually
# interleave mid-handler like they do under production load.
async def _yield_sleep(_seconds):
await asyncio.sleep(0)
monkeypatch.setattr(user_module, 'asyncio', SimpleNamespace(sleep=_yield_sleep))
client, reset_password, _ = await _create_client()
responses = await asyncio.gather(
*(client.post('/api/v1/user/reset-password', json=_payload(key='WRONG')) for _ in range(20))
)
status_counts: dict[int, int] = {}
for resp in responses:
status_counts[resp.status_code] = status_counts.get(resp.status_code, 0) + 1
assert status_counts == {403: 5, 429: 15}
reset_password.assert_not_awaited()