mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-09-30 21:36:47 +08:00
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:
@@ -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()
|
||||
Reference in New Issue
Block a user