mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-07-21 11:56:09 +00:00
test: format test suite
This commit is contained in:
+36
-36
@@ -58,45 +58,45 @@ from tests.factories.platform import (
|
||||
|
||||
__all__ = [
|
||||
# App
|
||||
"FakeApp",
|
||||
"fake_app",
|
||||
'FakeApp',
|
||||
'fake_app',
|
||||
# Message chains
|
||||
"text_chain",
|
||||
"group_text_chain",
|
||||
"mention_chain",
|
||||
"image_chain",
|
||||
'text_chain',
|
||||
'group_text_chain',
|
||||
'mention_chain',
|
||||
'image_chain',
|
||||
# Message events
|
||||
"friend_message_event",
|
||||
"group_message_event",
|
||||
'friend_message_event',
|
||||
'group_message_event',
|
||||
# Mock adapters
|
||||
"mock_adapter",
|
||||
'mock_adapter',
|
||||
# Queries
|
||||
"text_query",
|
||||
"group_text_query",
|
||||
"private_text_query",
|
||||
"command_query",
|
||||
"mention_query",
|
||||
"empty_query",
|
||||
"image_query",
|
||||
"file_query",
|
||||
"unsupported_query",
|
||||
"voice_query",
|
||||
"at_all_query",
|
||||
"query_with_session",
|
||||
"query_with_config",
|
||||
'text_query',
|
||||
'group_text_query',
|
||||
'private_text_query',
|
||||
'command_query',
|
||||
'mention_query',
|
||||
'empty_query',
|
||||
'image_query',
|
||||
'file_query',
|
||||
'unsupported_query',
|
||||
'voice_query',
|
||||
'at_all_query',
|
||||
'query_with_session',
|
||||
'query_with_config',
|
||||
# Provider
|
||||
"FakeProvider",
|
||||
"fake_provider",
|
||||
"fake_provider_pong",
|
||||
"fake_provider_timeout",
|
||||
"fake_provider_auth_error",
|
||||
"fake_provider_rate_limit",
|
||||
"fake_provider_malformed",
|
||||
"fake_model",
|
||||
'FakeProvider',
|
||||
'fake_provider',
|
||||
'fake_provider_pong',
|
||||
'fake_provider_timeout',
|
||||
'fake_provider_auth_error',
|
||||
'fake_provider_rate_limit',
|
||||
'fake_provider_malformed',
|
||||
'fake_model',
|
||||
# Platform
|
||||
"FakePlatform",
|
||||
"fake_platform",
|
||||
"fake_platform_with_streaming",
|
||||
"fake_platform_with_failure",
|
||||
"mock_platform_adapter",
|
||||
]
|
||||
'FakePlatform',
|
||||
'fake_platform',
|
||||
'fake_platform_with_streaming',
|
||||
'fake_platform_with_failure',
|
||||
'mock_platform_adapter',
|
||||
]
|
||||
|
||||
+83
-77
@@ -30,32 +30,36 @@ def _next_query_id() -> int:
|
||||
# ============== Message Chain Factories ==============
|
||||
|
||||
|
||||
def text_chain(text: str = "hello") -> platform_message.MessageChain:
|
||||
def text_chain(text: str = 'hello') -> platform_message.MessageChain:
|
||||
"""Create a simple text message chain."""
|
||||
return platform_message.MessageChain([
|
||||
platform_message.Plain(text=text),
|
||||
])
|
||||
return platform_message.MessageChain(
|
||||
[
|
||||
platform_message.Plain(text=text),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def group_text_chain(text: str = "hello") -> platform_message.MessageChain:
|
||||
def group_text_chain(text: str = 'hello') -> platform_message.MessageChain:
|
||||
"""Create a group text message chain (same as text_chain, context provided by event)."""
|
||||
return text_chain(text)
|
||||
|
||||
|
||||
def mention_chain(
|
||||
text: str = "hello",
|
||||
text: str = 'hello',
|
||||
target: typing.Union[int, str] = 12345,
|
||||
) -> platform_message.MessageChain:
|
||||
"""Create a message chain with @mention."""
|
||||
return platform_message.MessageChain([
|
||||
platform_message.At(target=target),
|
||||
platform_message.Plain(text=f" {text}"),
|
||||
])
|
||||
return platform_message.MessageChain(
|
||||
[
|
||||
platform_message.At(target=target),
|
||||
platform_message.Plain(text=f' {text}'),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def image_chain(
|
||||
text: str = "",
|
||||
url: str = "https://example.com/image.png",
|
||||
text: str = '',
|
||||
url: str = 'https://example.com/image.png',
|
||||
) -> platform_message.MessageChain:
|
||||
"""Create a message chain with an image."""
|
||||
components = []
|
||||
@@ -66,13 +70,15 @@ def image_chain(
|
||||
|
||||
|
||||
def command_chain(
|
||||
command: str = "help",
|
||||
prefix: str = "/",
|
||||
command: str = 'help',
|
||||
prefix: str = '/',
|
||||
) -> platform_message.MessageChain:
|
||||
"""Create a command message chain."""
|
||||
return platform_message.MessageChain([
|
||||
platform_message.Plain(text=f"{prefix}{command}"),
|
||||
])
|
||||
return platform_message.MessageChain(
|
||||
[
|
||||
platform_message.Plain(text=f'{prefix}{command}'),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# ============== Message Event Factories ==============
|
||||
@@ -81,7 +87,7 @@ def command_chain(
|
||||
def friend_message_event(
|
||||
message_chain: platform_message.MessageChain,
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
nickname: str = "TestUser",
|
||||
nickname: str = 'TestUser',
|
||||
) -> platform_events.FriendMessage:
|
||||
"""Create a friend (private) message event."""
|
||||
sender = platform_entities.Friend(
|
||||
@@ -90,7 +96,7 @@ def friend_message_event(
|
||||
remark=None,
|
||||
)
|
||||
return platform_events.FriendMessage(
|
||||
type="FriendMessage",
|
||||
type='FriendMessage',
|
||||
sender=sender,
|
||||
message_chain=message_chain,
|
||||
time=1609459200,
|
||||
@@ -100,9 +106,9 @@ def friend_message_event(
|
||||
def group_message_event(
|
||||
message_chain: platform_message.MessageChain,
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
sender_name: str = "TestUser",
|
||||
sender_name: str = 'TestUser',
|
||||
group_id: typing.Union[int, str] = 99999,
|
||||
group_name: str = "TestGroup",
|
||||
group_name: str = 'TestGroup',
|
||||
) -> platform_events.GroupMessage:
|
||||
"""Create a group message event."""
|
||||
group = platform_entities.Group(
|
||||
@@ -117,7 +123,7 @@ def group_message_event(
|
||||
group=group,
|
||||
)
|
||||
return platform_events.GroupMessage(
|
||||
type="GroupMessage",
|
||||
type='GroupMessage',
|
||||
sender=sender,
|
||||
message_chain=message_chain,
|
||||
time=1609459200,
|
||||
@@ -152,36 +158,36 @@ def _base_query(
|
||||
query_id = _next_query_id()
|
||||
|
||||
base_data = {
|
||||
"query_id": query_id,
|
||||
"launcher_type": launcher_type,
|
||||
"launcher_id": launcher_id,
|
||||
"sender_id": sender_id,
|
||||
"message_chain": message_chain,
|
||||
"message_event": message_event,
|
||||
"adapter": adapter,
|
||||
"pipeline_uuid": "test-pipeline-uuid",
|
||||
"bot_uuid": "test-bot-uuid",
|
||||
"pipeline_config": {
|
||||
"ai": {
|
||||
"runner": {"runner": "local-agent"},
|
||||
"local-agent": {
|
||||
"model": {"primary": "test-model-uuid", "fallbacks": []},
|
||||
"prompt": "test-prompt",
|
||||
'query_id': query_id,
|
||||
'launcher_type': launcher_type,
|
||||
'launcher_id': launcher_id,
|
||||
'sender_id': sender_id,
|
||||
'message_chain': message_chain,
|
||||
'message_event': message_event,
|
||||
'adapter': adapter,
|
||||
'pipeline_uuid': 'test-pipeline-uuid',
|
||||
'bot_uuid': 'test-bot-uuid',
|
||||
'pipeline_config': {
|
||||
'ai': {
|
||||
'runner': {'runner': 'local-agent'},
|
||||
'local-agent': {
|
||||
'model': {'primary': 'test-model-uuid', 'fallbacks': []},
|
||||
'prompt': 'test-prompt',
|
||||
},
|
||||
},
|
||||
"output": {"misc": {"at-sender": False, "quote-origin": False}},
|
||||
"trigger": {"misc": {"combine-quote-message": False}},
|
||||
'output': {'misc': {'at-sender': False, 'quote-origin': False}},
|
||||
'trigger': {'misc': {'combine-quote-message': False}},
|
||||
},
|
||||
"session": None,
|
||||
"prompt": None,
|
||||
"messages": [],
|
||||
"user_message": None,
|
||||
"use_funcs": [],
|
||||
"use_llm_model_uuid": None,
|
||||
"variables": {},
|
||||
"resp_messages": [],
|
||||
"resp_message_chain": None,
|
||||
"current_stage_name": None,
|
||||
'session': None,
|
||||
'prompt': None,
|
||||
'messages': [],
|
||||
'user_message': None,
|
||||
'use_funcs': [],
|
||||
'use_llm_model_uuid': None,
|
||||
'variables': {},
|
||||
'resp_messages': [],
|
||||
'resp_message_chain': None,
|
||||
'current_stage_name': None,
|
||||
}
|
||||
|
||||
# Apply overrides
|
||||
@@ -192,7 +198,7 @@ def _base_query(
|
||||
|
||||
|
||||
def text_query(
|
||||
text: str = "hello",
|
||||
text: str = 'hello',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
**overrides,
|
||||
) -> pipeline_query.Query:
|
||||
@@ -212,7 +218,7 @@ def text_query(
|
||||
|
||||
|
||||
def private_text_query(
|
||||
text: str = "hello",
|
||||
text: str = 'hello',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
**overrides,
|
||||
) -> pipeline_query.Query:
|
||||
@@ -221,7 +227,7 @@ def private_text_query(
|
||||
|
||||
|
||||
def group_text_query(
|
||||
text: str = "hello",
|
||||
text: str = 'hello',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
group_id: typing.Union[int, str] = 99999,
|
||||
**overrides,
|
||||
@@ -242,8 +248,8 @@ def group_text_query(
|
||||
|
||||
|
||||
def command_query(
|
||||
command: str = "help",
|
||||
prefix: str = "/",
|
||||
command: str = 'help',
|
||||
prefix: str = '/',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
**overrides,
|
||||
) -> pipeline_query.Query:
|
||||
@@ -263,7 +269,7 @@ def command_query(
|
||||
|
||||
|
||||
def mention_query(
|
||||
text: str = "hello",
|
||||
text: str = 'hello',
|
||||
target: typing.Union[int, str] = 12345,
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
group_id: typing.Union[int, str] = 99999,
|
||||
@@ -301,8 +307,8 @@ def empty_query(**overrides) -> pipeline_query.Query:
|
||||
|
||||
|
||||
def image_query(
|
||||
text: str = "",
|
||||
url: str = "https://example.com/image.png",
|
||||
text: str = '',
|
||||
url: str = 'https://example.com/image.png',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
**overrides,
|
||||
) -> pipeline_query.Query:
|
||||
@@ -322,9 +328,9 @@ def image_query(
|
||||
|
||||
|
||||
def file_query(
|
||||
url: str = "https://example.com/document.pdf",
|
||||
name: str = "document.pdf",
|
||||
text: str = "",
|
||||
url: str = 'https://example.com/document.pdf',
|
||||
name: str = 'document.pdf',
|
||||
text: str = '',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
**overrides,
|
||||
) -> pipeline_query.Query:
|
||||
@@ -348,8 +354,8 @@ def file_query(
|
||||
|
||||
|
||||
def unsupported_query(
|
||||
unsupported_type: str = "CustomComponent",
|
||||
text: str = "",
|
||||
unsupported_type: str = 'CustomComponent',
|
||||
text: str = '',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
**overrides,
|
||||
) -> pipeline_query.Query:
|
||||
@@ -358,7 +364,7 @@ def unsupported_query(
|
||||
if text:
|
||||
components.append(platform_message.Plain(text=text))
|
||||
# Use Unknown component for unsupported types
|
||||
components.append(platform_message.Unknown(text=f"Unsupported: {unsupported_type}"))
|
||||
components.append(platform_message.Unknown(text=f'Unsupported: {unsupported_type}'))
|
||||
chain = platform_message.MessageChain(components)
|
||||
event = friend_message_event(chain, sender_id)
|
||||
adapter = mock_adapter()
|
||||
@@ -374,7 +380,7 @@ def unsupported_query(
|
||||
|
||||
|
||||
def query_with_session(
|
||||
text: str = "hello",
|
||||
text: str = 'hello',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
session: provider_session.Session = None,
|
||||
**overrides,
|
||||
@@ -389,7 +395,7 @@ def query_with_session(
|
||||
launcher_type=provider_session.LauncherTypes.PERSON,
|
||||
launcher_id=sender_id,
|
||||
sender_id=sender_id,
|
||||
use_prompt_name="default",
|
||||
use_prompt_name='default',
|
||||
using_conversation=None,
|
||||
conversations=[],
|
||||
)
|
||||
@@ -398,7 +404,7 @@ def query_with_session(
|
||||
|
||||
|
||||
def query_with_config(
|
||||
text: str = "hello",
|
||||
text: str = 'hello',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
pipeline_config: dict = None,
|
||||
**overrides,
|
||||
@@ -410,22 +416,22 @@ def query_with_config(
|
||||
"""
|
||||
if pipeline_config is None:
|
||||
pipeline_config = {
|
||||
"ai": {
|
||||
"runner": {"runner": "local-agent"},
|
||||
"local-agent": {
|
||||
"model": {"primary": "test-model-uuid", "fallbacks": []},
|
||||
"prompt": "test-prompt",
|
||||
'ai': {
|
||||
'runner': {'runner': 'local-agent'},
|
||||
'local-agent': {
|
||||
'model': {'primary': 'test-model-uuid', 'fallbacks': []},
|
||||
'prompt': 'test-prompt',
|
||||
},
|
||||
},
|
||||
"output": {"misc": {"at-sender": False, "quote-origin": False}},
|
||||
"trigger": {"misc": {"combine-quote-message": False}},
|
||||
'output': {'misc': {'at-sender': False, 'quote-origin': False}},
|
||||
'trigger': {'misc': {'combine-quote-message': False}},
|
||||
}
|
||||
|
||||
return text_query(text, sender_id, pipeline_config=pipeline_config, **overrides)
|
||||
|
||||
|
||||
def voice_query(
|
||||
url: str = "https://example.com/audio.mp3",
|
||||
url: str = 'https://example.com/audio.mp3',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
**overrides,
|
||||
) -> pipeline_query.Query:
|
||||
@@ -448,7 +454,7 @@ def voice_query(
|
||||
|
||||
|
||||
def at_all_query(
|
||||
text: str = "hello",
|
||||
text: str = 'hello',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
group_id: typing.Union[int, str] = 99999,
|
||||
**overrides,
|
||||
@@ -456,7 +462,7 @@ def at_all_query(
|
||||
"""Create a group query with @All mention."""
|
||||
components = [
|
||||
platform_message.AtAll(),
|
||||
platform_message.Plain(text=f" {text}"),
|
||||
platform_message.Plain(text=f' {text}'),
|
||||
]
|
||||
chain = platform_message.MessageChain(components)
|
||||
event = group_message_event(chain, sender_id, group_id=group_id)
|
||||
@@ -469,4 +475,4 @@ def at_all_query(
|
||||
sender_id=sender_id,
|
||||
adapter=adapter,
|
||||
**overrides,
|
||||
)
|
||||
)
|
||||
|
||||
+52
-46
@@ -33,7 +33,7 @@ class FakePlatform:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
bot_account_id: str = "test-bot",
|
||||
bot_account_id: str = 'test-bot',
|
||||
stream_output_supported: bool = False,
|
||||
raise_error: Exception = None,
|
||||
):
|
||||
@@ -48,16 +48,16 @@ class FakePlatform:
|
||||
# Registered listeners
|
||||
self._listeners: dict = {}
|
||||
|
||||
def raises(self, error: Exception) -> "FakePlatform":
|
||||
def raises(self, error: Exception) -> 'FakePlatform':
|
||||
"""Configure platform to raise an error on send."""
|
||||
self._raise_error = error
|
||||
return self
|
||||
|
||||
def send_failure(self) -> "FakePlatform":
|
||||
def send_failure(self) -> 'FakePlatform':
|
||||
"""Configure platform to simulate send failure."""
|
||||
return self.raises(Exception("Platform send failure"))
|
||||
return self.raises(Exception('Platform send failure'))
|
||||
|
||||
def supports_streaming(self, supported: bool = True) -> "FakePlatform":
|
||||
def supports_streaming(self, supported: bool = True) -> 'FakePlatform':
|
||||
"""Configure whether streaming output is supported."""
|
||||
self._stream_output_supported = supported
|
||||
return self
|
||||
@@ -89,7 +89,7 @@ class FakePlatform:
|
||||
self,
|
||||
text: str,
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
nickname: str = "TestUser",
|
||||
nickname: str = 'TestUser',
|
||||
) -> platform_events.FriendMessage:
|
||||
"""Create an inbound friend (private) message event."""
|
||||
sender = platform_entities.Friend(
|
||||
@@ -97,11 +97,13 @@ class FakePlatform:
|
||||
nickname=nickname,
|
||||
remark=None,
|
||||
)
|
||||
chain = platform_message.MessageChain([
|
||||
platform_message.Plain(text=text),
|
||||
])
|
||||
chain = platform_message.MessageChain(
|
||||
[
|
||||
platform_message.Plain(text=text),
|
||||
]
|
||||
)
|
||||
return platform_events.FriendMessage(
|
||||
type="FriendMessage",
|
||||
type='FriendMessage',
|
||||
sender=sender,
|
||||
message_chain=chain,
|
||||
time=1609459200,
|
||||
@@ -111,9 +113,9 @@ class FakePlatform:
|
||||
self,
|
||||
text: str,
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
sender_name: str = "TestUser",
|
||||
sender_name: str = 'TestUser',
|
||||
group_id: typing.Union[int, str] = 99999,
|
||||
group_name: str = "TestGroup",
|
||||
group_name: str = 'TestGroup',
|
||||
mention_bot: bool = False,
|
||||
) -> platform_events.GroupMessage:
|
||||
"""Create an inbound group message event.
|
||||
@@ -142,12 +144,12 @@ class FakePlatform:
|
||||
components = []
|
||||
if mention_bot:
|
||||
components.append(platform_message.At(target=self.bot_account_id))
|
||||
components.append(platform_message.Plain(text=" "))
|
||||
components.append(platform_message.Plain(text=' '))
|
||||
components.append(platform_message.Plain(text=text))
|
||||
|
||||
chain = platform_message.MessageChain(components)
|
||||
return platform_events.GroupMessage(
|
||||
type="GroupMessage",
|
||||
type='GroupMessage',
|
||||
sender=sender,
|
||||
message_chain=chain,
|
||||
time=1609459200,
|
||||
@@ -155,8 +157,8 @@ class FakePlatform:
|
||||
|
||||
def create_image_message(
|
||||
self,
|
||||
url: str = "https://example.com/image.png",
|
||||
text: str = "",
|
||||
url: str = 'https://example.com/image.png',
|
||||
text: str = '',
|
||||
sender_id: typing.Union[int, str] = 12345,
|
||||
is_group: bool = False,
|
||||
group_id: typing.Union[int, str] = 99999,
|
||||
@@ -169,12 +171,12 @@ class FakePlatform:
|
||||
chain = platform_message.MessageChain(components)
|
||||
|
||||
if is_group:
|
||||
return self.create_group_message("", sender_id, group_id=group_id)
|
||||
return self.create_group_message('', sender_id, group_id=group_id)
|
||||
# Replace chain
|
||||
else:
|
||||
sender = platform_entities.Friend(id=sender_id, nickname="TestUser", remark=None)
|
||||
sender = platform_entities.Friend(id=sender_id, nickname='TestUser', remark=None)
|
||||
return platform_events.FriendMessage(
|
||||
type="FriendMessage",
|
||||
type='FriendMessage',
|
||||
sender=sender,
|
||||
message_chain=chain,
|
||||
time=1609459200,
|
||||
@@ -192,12 +194,14 @@ class FakePlatform:
|
||||
if self._raise_error:
|
||||
raise self._raise_error
|
||||
|
||||
self._outbound_messages.append({
|
||||
"type": "send",
|
||||
"target_type": target_type,
|
||||
"target_id": target_id,
|
||||
"message": message,
|
||||
})
|
||||
self._outbound_messages.append(
|
||||
{
|
||||
'type': 'send',
|
||||
'target_type': target_type,
|
||||
'target_id': target_id,
|
||||
'message': message,
|
||||
}
|
||||
)
|
||||
|
||||
async def reply_message(
|
||||
self,
|
||||
@@ -209,13 +213,15 @@ class FakePlatform:
|
||||
if self._raise_error:
|
||||
raise self._raise_error
|
||||
|
||||
self._outbound_messages.append({
|
||||
"type": "reply",
|
||||
"source_type": message_source.type,
|
||||
"source": message_source,
|
||||
"message": message,
|
||||
"quote_origin": quote_origin,
|
||||
})
|
||||
self._outbound_messages.append(
|
||||
{
|
||||
'type': 'reply',
|
||||
'source_type': message_source.type,
|
||||
'source': message_source,
|
||||
'message': message,
|
||||
'quote_origin': quote_origin,
|
||||
}
|
||||
)
|
||||
|
||||
async def reply_message_chunk(
|
||||
self,
|
||||
@@ -229,15 +235,17 @@ class FakePlatform:
|
||||
if self._raise_error:
|
||||
raise self._raise_error
|
||||
|
||||
self._outbound_chunks.append({
|
||||
"type": "reply_chunk",
|
||||
"source_type": message_source.type,
|
||||
"source": message_source,
|
||||
"bot_message": bot_message,
|
||||
"message": message,
|
||||
"quote_origin": quote_origin,
|
||||
"is_final": is_final,
|
||||
})
|
||||
self._outbound_chunks.append(
|
||||
{
|
||||
'type': 'reply_chunk',
|
||||
'source_type': message_source.type,
|
||||
'source': message_source,
|
||||
'bot_message': bot_message,
|
||||
'message': message,
|
||||
'quote_origin': quote_origin,
|
||||
'is_final': is_final,
|
||||
}
|
||||
)
|
||||
|
||||
async def is_stream_output_supported(self) -> bool:
|
||||
"""Return whether streaming output is supported."""
|
||||
@@ -295,7 +303,7 @@ class FakePlatform:
|
||||
|
||||
|
||||
def fake_platform(
|
||||
bot_account_id: str = "test-bot",
|
||||
bot_account_id: str = 'test-bot',
|
||||
stream_output_supported: bool = False,
|
||||
) -> FakePlatform:
|
||||
"""Create a FakePlatform instance."""
|
||||
@@ -328,9 +336,7 @@ def mock_platform_adapter(platform: FakePlatform = None) -> Mock:
|
||||
adapter.reply_message = AsyncMock(side_effect=platform.reply_message)
|
||||
adapter.reply_message_chunk = AsyncMock(side_effect=platform.reply_message_chunk)
|
||||
adapter.send_message = AsyncMock(side_effect=platform.send_message)
|
||||
adapter.is_stream_output_supported = AsyncMock(
|
||||
return_value=platform._stream_output_supported
|
||||
)
|
||||
adapter.is_stream_output_supported = AsyncMock(return_value=platform._stream_output_supported)
|
||||
adapter._fake_platform = platform # Store for assertions
|
||||
|
||||
return adapter
|
||||
return adapter
|
||||
|
||||
+42
-38
@@ -27,51 +27,51 @@ class FakeProvider:
|
||||
Does not require API keys.
|
||||
"""
|
||||
|
||||
PONG_RESPONSE = "LANGBOT_FAKE_PONG"
|
||||
PONG_RESPONSE = 'LANGBOT_FAKE_PONG'
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
default_response: str = "fake response",
|
||||
default_response: str = 'fake response',
|
||||
streaming_chunks: list[str] = None,
|
||||
raise_error: Exception = None,
|
||||
captured_requests: list = None,
|
||||
):
|
||||
self._default_response = default_response
|
||||
self._streaming_chunks = streaming_chunks or ["fake ", "response"]
|
||||
self._streaming_chunks = streaming_chunks or ['fake ', 'response']
|
||||
self._raise_error = raise_error
|
||||
self._captured_requests = captured_requests if captured_requests is not None else []
|
||||
|
||||
def returns(self, text: str) -> "FakeProvider":
|
||||
def returns(self, text: str) -> 'FakeProvider':
|
||||
"""Configure provider to return a specific text response."""
|
||||
self._default_response = text
|
||||
self._streaming_chunks = [text]
|
||||
return self
|
||||
|
||||
def returns_streaming(self, chunks: list[str]) -> "FakeProvider":
|
||||
def returns_streaming(self, chunks: list[str]) -> 'FakeProvider':
|
||||
"""Configure provider to return streaming chunks."""
|
||||
self._streaming_chunks = chunks
|
||||
self._default_response = "".join(chunks)
|
||||
self._default_response = ''.join(chunks)
|
||||
return self
|
||||
|
||||
def raises(self, error: Exception) -> "FakeProvider":
|
||||
def raises(self, error: Exception) -> 'FakeProvider':
|
||||
"""Configure provider to raise an error."""
|
||||
self._raise_error = error
|
||||
return self
|
||||
|
||||
def timeout(self) -> "FakeProvider":
|
||||
def timeout(self) -> 'FakeProvider':
|
||||
"""Configure provider to simulate timeout."""
|
||||
return self.raises(TimeoutError("Provider timeout"))
|
||||
return self.raises(TimeoutError('Provider timeout'))
|
||||
|
||||
def auth_error(self) -> "FakeProvider":
|
||||
def auth_error(self) -> 'FakeProvider':
|
||||
"""Configure provider to simulate auth error."""
|
||||
return self.raises(Exception("Invalid API key"))
|
||||
return self.raises(Exception('Invalid API key'))
|
||||
|
||||
def rate_limit(self) -> "FakeProvider":
|
||||
def rate_limit(self) -> 'FakeProvider':
|
||||
"""Configure provider to simulate rate limit."""
|
||||
return self.raises(Exception("Rate limit exceeded"))
|
||||
return self.raises(Exception('Rate limit exceeded'))
|
||||
|
||||
def malformed(self) -> "FakeProvider":
|
||||
def malformed(self) -> 'FakeProvider':
|
||||
"""Configure provider to simulate malformed response."""
|
||||
self._default_response = None
|
||||
return self
|
||||
@@ -87,7 +87,7 @@ class FakeProvider:
|
||||
def _create_message(self, content: str) -> provider_message.Message:
|
||||
"""Create a provider message from text content."""
|
||||
return provider_message.Message(
|
||||
role="assistant",
|
||||
role='assistant',
|
||||
content=content,
|
||||
)
|
||||
|
||||
@@ -99,7 +99,7 @@ class FakeProvider:
|
||||
) -> provider_message.MessageChunk:
|
||||
"""Create a provider message chunk."""
|
||||
return provider_message.MessageChunk(
|
||||
role="assistant",
|
||||
role='assistant',
|
||||
content=content,
|
||||
is_final=is_final,
|
||||
msg_sequence=msg_sequence,
|
||||
@@ -116,13 +116,15 @@ class FakeProvider:
|
||||
) -> provider_message.Message:
|
||||
"""Simulate non-streaming LLM invocation."""
|
||||
# Capture request for assertions
|
||||
self._captured_requests.append({
|
||||
"query_id": query.query_id if query else None,
|
||||
"model": model.model_entity.name if model and hasattr(model, 'model_entity') else None,
|
||||
"messages": messages,
|
||||
"funcs": funcs,
|
||||
"extra_args": extra_args,
|
||||
})
|
||||
self._captured_requests.append(
|
||||
{
|
||||
'query_id': query.query_id if query else None,
|
||||
'model': model.model_entity.name if model and hasattr(model, 'model_entity') else None,
|
||||
'messages': messages,
|
||||
'funcs': funcs,
|
||||
'extra_args': extra_args,
|
||||
}
|
||||
)
|
||||
|
||||
# Simulate error if configured
|
||||
if self._raise_error:
|
||||
@@ -131,7 +133,7 @@ class FakeProvider:
|
||||
# Return response
|
||||
if self._default_response is None:
|
||||
# Malformed response
|
||||
return provider_message.Message(role="assistant", content=None)
|
||||
return provider_message.Message(role='assistant', content=None)
|
||||
|
||||
return self._create_message(self._default_response)
|
||||
|
||||
@@ -146,14 +148,16 @@ class FakeProvider:
|
||||
) -> typing.AsyncGenerator[provider_message.MessageChunk, None]:
|
||||
"""Simulate streaming LLM invocation."""
|
||||
# Capture request for assertions
|
||||
self._captured_requests.append({
|
||||
"query_id": query.query_id if query else None,
|
||||
"model": model.model_entity.name if model and hasattr(model, 'model_entity') else None,
|
||||
"messages": messages,
|
||||
"funcs": funcs,
|
||||
"extra_args": extra_args,
|
||||
"streaming": True,
|
||||
})
|
||||
self._captured_requests.append(
|
||||
{
|
||||
'query_id': query.query_id if query else None,
|
||||
'model': model.model_entity.name if model and hasattr(model, 'model_entity') else None,
|
||||
'messages': messages,
|
||||
'funcs': funcs,
|
||||
'extra_args': extra_args,
|
||||
'streaming': True,
|
||||
}
|
||||
)
|
||||
|
||||
# Simulate error if configured
|
||||
if self._raise_error:
|
||||
@@ -161,12 +165,12 @@ class FakeProvider:
|
||||
|
||||
# Yield chunks
|
||||
for i, chunk in enumerate(self._streaming_chunks):
|
||||
is_final = (i == len(self._streaming_chunks) - 1)
|
||||
is_final = i == len(self._streaming_chunks) - 1
|
||||
yield self._create_chunk(chunk, is_final=is_final, msg_sequence=i)
|
||||
|
||||
|
||||
def fake_provider(
|
||||
default_response: str = "fake response",
|
||||
default_response: str = 'fake response',
|
||||
) -> FakeProvider:
|
||||
"""Create a FakeProvider with optional default response."""
|
||||
return FakeProvider(default_response=default_response)
|
||||
@@ -202,8 +206,8 @@ def fake_provider_malformed() -> FakeProvider:
|
||||
|
||||
def fake_model(
|
||||
*,
|
||||
uuid: str = "test-model-uuid",
|
||||
name: str = "test-model",
|
||||
uuid: str = 'test-model-uuid',
|
||||
name: str = 'test-model',
|
||||
abilities: list[str] = None,
|
||||
provider: FakeProvider = None,
|
||||
) -> Mock:
|
||||
@@ -212,7 +216,7 @@ def fake_model(
|
||||
model.model_entity = Mock()
|
||||
model.model_entity.uuid = uuid
|
||||
model.model_entity.name = name
|
||||
model.model_entity.abilities = abilities or ["func_call", "vision"]
|
||||
model.model_entity.abilities = abilities or ['func_call', 'vision']
|
||||
model.model_entity.extra_args = {}
|
||||
|
||||
# Attach fake provider
|
||||
@@ -221,4 +225,4 @@ def fake_model(
|
||||
|
||||
model.provider = provider
|
||||
|
||||
return model
|
||||
return model
|
||||
|
||||
Reference in New Issue
Block a user