mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
test: mass-coverage push — ~1000+ new unit tests
Coverage improvements: - agent/tool/core: 150+ tests (FunctionTool, ToolSet, ToolSchema, run_context, messages, registry) - provider: 152 tests (entities, Provider ABC, register, manager) - star module: 58 tests (Context, PluginManager, StarHandler) - platform: 71 tests (Platform ABC, manager, MessageSession) - pipeline: 130 tests (process/respond stages, scheduler, bootstrap, conversation_mgr) - kb+skills: 96 tests (mgr, helpers, chunkers, skill_manager) - config+builtins: 92 tests (astrbot_config, config_mgr, default, LTM, star main) - cron+db+event_bus: 89 tests (cron events/manager, PO models, BaseDatabase, EventBus) - dashboard routes: expanded smoke tests for all route modules
This commit is contained in:
@@ -0,0 +1,280 @@
|
||||
"""Unit tests for astrbot.core.agent.message: Message, CheckpointData, CheckpointMessageSegment."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from astrbot.core.agent.message import (
|
||||
AssistantMessageSegment,
|
||||
CheckpointData,
|
||||
CheckpointMessageSegment,
|
||||
ContentPart,
|
||||
ImageURLPart,
|
||||
Message,
|
||||
SystemMessageSegment,
|
||||
TextPart,
|
||||
ThinkPart,
|
||||
ToolCall,
|
||||
ToolCallMessageSegment,
|
||||
UserMessageSegment,
|
||||
bind_checkpoint_messages,
|
||||
dump_messages_with_checkpoints,
|
||||
get_checkpoint_id,
|
||||
is_checkpoint_message,
|
||||
strip_checkpoint_messages,
|
||||
)
|
||||
|
||||
|
||||
class TestMessageConstruction:
|
||||
"""Message base model construction and validation."""
|
||||
|
||||
def test_minimal_user_message(self):
|
||||
"""A user message requires role and content."""
|
||||
msg = Message(role="user", content="hello")
|
||||
assert msg.role == "user"
|
||||
assert msg.content == "hello"
|
||||
assert msg.tool_calls is None
|
||||
assert msg.tool_call_id is None
|
||||
|
||||
def test_assistant_message_with_content(self):
|
||||
"""An assistant message with text content."""
|
||||
msg = Message(role="assistant", content="I am an assistant")
|
||||
assert msg.role == "assistant"
|
||||
assert msg.content == "I am an assistant"
|
||||
|
||||
def test_assistant_message_with_tool_calls_no_content(self):
|
||||
"""Assistant messages with tool_calls may have content=None."""
|
||||
tc = ToolCall(id="call_1", function=ToolCall.FunctionBody(name="f", arguments="{}"))
|
||||
msg = Message(role="assistant", content=None, tool_calls=[tc])
|
||||
assert msg.role == "assistant"
|
||||
assert msg.content is None
|
||||
assert len(msg.tool_calls or []) == 1
|
||||
|
||||
def test_tool_message(self):
|
||||
"""A tool message with role='tool'."""
|
||||
msg = Message(role="tool", content="tool result", tool_call_id="call_1")
|
||||
assert msg.role == "tool"
|
||||
assert msg.tool_call_id == "call_1"
|
||||
|
||||
def test_system_message(self):
|
||||
"""A system message."""
|
||||
msg = Message(role="system", content="You are a helpful assistant")
|
||||
assert msg.role == "system"
|
||||
|
||||
def test_missing_content_raises_for_user(self):
|
||||
"""User/System/Tool messages must have content."""
|
||||
with pytest.raises(ValidationError, match="content is required"):
|
||||
Message(role="user", content=None)
|
||||
|
||||
def test_missing_content_raises_for_system(self):
|
||||
"""System messages must have content."""
|
||||
with pytest.raises(ValidationError, match="content is required"):
|
||||
Message(role="system", content=None)
|
||||
|
||||
def test_invalid_role_raises(self):
|
||||
"""An invalid role is rejected."""
|
||||
with pytest.raises(ValidationError):
|
||||
Message(role="invalid_role", content="hi")
|
||||
|
||||
|
||||
class TestCheckpointMessage:
|
||||
"""CheckpointData and role='_checkpoint' messages."""
|
||||
|
||||
def test_checkpoint_data_construction(self):
|
||||
"""CheckpointData can be constructed with an id."""
|
||||
cp = CheckpointData(id="cp_1")
|
||||
assert cp.id == "cp_1"
|
||||
|
||||
def test_checkpoint_message_valid(self):
|
||||
"""A valid checkpoint message has role _checkpoint and CheckpointData content."""
|
||||
cp = CheckpointData(id="cp_1")
|
||||
msg = Message(role="_checkpoint", content=cp)
|
||||
assert msg.role == "_checkpoint"
|
||||
assert isinstance(msg.content, CheckpointData)
|
||||
|
||||
def test_checkpoint_message_string_content_raises(self):
|
||||
"""Checkpoint messages must use CheckpointData, not plain strings."""
|
||||
with pytest.raises(ValidationError):
|
||||
Message(role="_checkpoint", content="not a checkpoint")
|
||||
|
||||
def test_checkpoint_data_in_non_checkpoint_role_raises(self):
|
||||
"""CheckpointData in a non-checkpoint role is rejected."""
|
||||
cp = CheckpointData(id="cp_1")
|
||||
with pytest.raises(ValidationError, match="CheckpointData is only allowed"):
|
||||
Message(role="user", content=cp)
|
||||
|
||||
def test_is_checkpoint_message_detection(self):
|
||||
"""is_checkpoint_message correctly identifies checkpoint messages."""
|
||||
cp = CheckpointData(id="cp_1")
|
||||
msg = Message(role="_checkpoint", content=cp)
|
||||
assert is_checkpoint_message(msg) is True
|
||||
assert is_checkpoint_message(Message(role="user", content="hi")) is False
|
||||
|
||||
def test_is_checkpoint_message_dict(self):
|
||||
"""is_checkpoint_message works with dicts."""
|
||||
assert is_checkpoint_message({"role": "_checkpoint"}) is True
|
||||
assert is_checkpoint_message({"role": "user"}) is False
|
||||
|
||||
def test_get_checkpoint_id(self):
|
||||
"""get_checkpoint_id returns the id from a checkpoint message."""
|
||||
cp = CheckpointData(id="cp_42")
|
||||
msg = Message(role="_checkpoint", content=cp)
|
||||
assert get_checkpoint_id(msg) == "cp_42"
|
||||
|
||||
def test_get_checkpoint_id_none_for_non_checkpoint(self):
|
||||
"""get_checkpoint_id returns None for non-checkpoint messages."""
|
||||
msg = Message(role="user", content="hi")
|
||||
assert get_checkpoint_id(msg) is None
|
||||
|
||||
def test_strip_checkpoint_messages(self):
|
||||
"""strip_checkpoint_messages removes checkpoint entries."""
|
||||
history = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "_checkpoint", "content": {"id": "cp_1"}},
|
||||
{"role": "assistant", "content": "hello"},
|
||||
]
|
||||
cleaned = strip_checkpoint_messages(history)
|
||||
assert len(cleaned) == 2
|
||||
assert all(m["role"] != "_checkpoint" for m in cleaned)
|
||||
|
||||
def test_bind_and_dump_checkpoints(self):
|
||||
"""dump_messages_with_checkpoints reinserts bound checkpoints after dump."""
|
||||
cp = CheckpointData(id="cp_99")
|
||||
msg = Message(role="assistant", content="hi")
|
||||
msg._checkpoint_after = cp
|
||||
dumped = dump_messages_with_checkpoints([msg])
|
||||
assert len(dumped) == 2
|
||||
assert dumped[0]["role"] == "assistant"
|
||||
assert dumped[1]["role"] == "_checkpoint"
|
||||
assert dumped[1]["content"]["id"] == "cp_99"
|
||||
|
||||
def test_bind_checkpoint_messages_roundtrip(self):
|
||||
"""bind_checkpoint_messages binds checkpoints to prior messages."""
|
||||
history = [
|
||||
{"role": "user", "content": "q"},
|
||||
{"role": "assistant", "content": "a"},
|
||||
{"role": "_checkpoint", "content": {"id": "cp_1"}},
|
||||
]
|
||||
messages = bind_checkpoint_messages(history)
|
||||
assert len(messages) == 2
|
||||
assert messages[1]._checkpoint_after is not None
|
||||
assert messages[1]._checkpoint_after.id == "cp_1"
|
||||
|
||||
|
||||
class TestMessageSegments:
|
||||
"""Typed message segment subclasses."""
|
||||
|
||||
def test_assistant_message_segment(self):
|
||||
"""AssistantMessageSegment has role fixed to 'assistant'."""
|
||||
msg = AssistantMessageSegment(content="hello")
|
||||
assert msg.role == "assistant"
|
||||
|
||||
def test_user_message_segment(self):
|
||||
"""UserMessageSegment has role fixed to 'user'."""
|
||||
msg = UserMessageSegment(content="hello")
|
||||
assert msg.role == "user"
|
||||
|
||||
def test_system_message_segment(self):
|
||||
"""SystemMessageSegment has role fixed to 'system'."""
|
||||
msg = SystemMessageSegment(content="beep")
|
||||
assert msg.role == "system"
|
||||
|
||||
def test_tool_call_message_segment(self):
|
||||
"""ToolCallMessageSegment has role fixed to 'tool'."""
|
||||
msg = ToolCallMessageSegment(content="result", tool_call_id="c1")
|
||||
assert msg.role == "tool"
|
||||
|
||||
def test_checkpoint_message_segment(self):
|
||||
"""CheckpointMessageSegment has role fixed to '_checkpoint' and optional CheckpointData content."""
|
||||
cp = CheckpointData(id="cp_1")
|
||||
msg = CheckpointMessageSegment(content=cp)
|
||||
assert msg.role == "_checkpoint"
|
||||
assert msg.content.id == "cp_1"
|
||||
|
||||
|
||||
class TestContentParts:
|
||||
"""Content part subclasses."""
|
||||
|
||||
def test_text_part(self):
|
||||
"""TextPart holds text content."""
|
||||
tp = TextPart(text="Hello, world!")
|
||||
assert tp.type == "text"
|
||||
assert tp.text == "Hello, world!"
|
||||
|
||||
def test_think_part(self):
|
||||
"""ThinkPart holds think content."""
|
||||
tp = ThinkPart(think="I need to think about this.")
|
||||
assert tp.type == "think"
|
||||
assert tp.think == "I need to think about this."
|
||||
assert tp.encrypted is None
|
||||
|
||||
def test_think_part_merge(self):
|
||||
"""merge_in_place appends think content."""
|
||||
t1 = ThinkPart(think="First ")
|
||||
t2 = ThinkPart(think="Second")
|
||||
assert t1.merge_in_place(t2) is True
|
||||
assert t1.think == "First Second"
|
||||
|
||||
def test_think_part_merge_non_think_returns_false(self):
|
||||
"""merge_in_place returns False when other is not a ThinkPart."""
|
||||
t1 = ThinkPart(think="A")
|
||||
assert t1.merge_in_place("not a think part") is False
|
||||
|
||||
def test_think_part_merge_encrypted_returns_false(self):
|
||||
"""merge_in_place returns False when self is encrypted."""
|
||||
t1 = ThinkPart(think="A", encrypted="sig1")
|
||||
t2 = ThinkPart(think="B")
|
||||
assert t1.merge_in_place(t2) is False
|
||||
|
||||
def test_image_url_part(self):
|
||||
"""ImageURLPart holds an image URL."""
|
||||
part = ImageURLPart(image_url=ImageURLPart.ImageURL(url="http://example.com/img.jpg"))
|
||||
assert part.type == "image_url"
|
||||
assert part.image_url.url == "http://example.com/img.jpg"
|
||||
|
||||
def test_image_url_with_id(self):
|
||||
"""ImageURLPart can include an id."""
|
||||
part = ImageURLPart(
|
||||
image_url=ImageURLPart.ImageURL(url="http://example.com/img.jpg", id="img_1")
|
||||
)
|
||||
assert part.image_url.id == "img_1"
|
||||
|
||||
|
||||
class TestToolCall:
|
||||
"""ToolCall construction and serialization."""
|
||||
|
||||
def test_tool_call_minimal(self):
|
||||
"""ToolCall with id and function body."""
|
||||
tc = ToolCall(id="call_1", function=ToolCall.FunctionBody(name="f", arguments="{}"))
|
||||
assert tc.id == "call_1"
|
||||
assert tc.function.name == "f"
|
||||
assert tc.function.arguments == "{}"
|
||||
|
||||
def test_tool_call_extra_content(self):
|
||||
"""ToolCall with extra_content is serialized."""
|
||||
tc = ToolCall(
|
||||
id="call_2",
|
||||
function=ToolCall.FunctionBody(name="g", arguments='{"x": 1}'),
|
||||
extra_content={"meta": "data"},
|
||||
)
|
||||
dumped = tc.model_dump()
|
||||
assert dumped["extra_content"] == {"meta": "data"}
|
||||
|
||||
def test_tool_call_extra_content_none_omitted(self):
|
||||
"""ToolCall with extra_content=None omits the field in serialization."""
|
||||
tc = ToolCall(id="call_3", function=ToolCall.FunctionBody(name="h", arguments="{}"))
|
||||
dumped = tc.model_dump()
|
||||
assert "extra_content" not in dumped
|
||||
|
||||
def test_message_serialization_omits_tool_calls_when_none(self):
|
||||
"""Message.model_dump omits tool_calls when None."""
|
||||
msg = Message(role="assistant", content="hi")
|
||||
dumped = msg.model_dump()
|
||||
assert "tool_calls" not in dumped
|
||||
|
||||
def test_message_serialization_omits_tool_call_id_when_none(self):
|
||||
"""Message.model_dump omits tool_call_id when None."""
|
||||
msg = Message(role="user", content="hi")
|
||||
dumped = msg.model_dump()
|
||||
assert "tool_call_id" not in dumped
|
||||
@@ -0,0 +1,323 @@
|
||||
"""Mock-based unit tests for AstrBotConfig."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.config.astrbot_config import AstrBotConfig, RateLimitStrategy
|
||||
|
||||
|
||||
class TestRateLimitStrategy:
|
||||
"""Tests for the RateLimitStrategy enum."""
|
||||
|
||||
def test_stall_member(self):
|
||||
assert RateLimitStrategy.STALL.value == "stall"
|
||||
|
||||
def test_discard_member(self):
|
||||
assert RateLimitStrategy.DISCARD.value == "discard"
|
||||
|
||||
|
||||
class TestAstrBotConfigConstruction:
|
||||
"""Construction and file-existence branches."""
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=False)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"version": 1}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_init_creates_file_when_missing(
|
||||
self, mock_json_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""When the file does not exist, __init__ writes the default config."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"version": 1},
|
||||
)
|
||||
mock_json_dump.assert_called_once()
|
||||
assert config["version"] == 1
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"key_a": "value_a", "key_b": 42}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_init_loads_existing_file(
|
||||
self, mock_json_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""When the file exists, __init__ loads its contents."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"key_a": "default", "key_b": 0},
|
||||
)
|
||||
assert config["key_a"] == "value_a"
|
||||
assert config["key_b"] == 42
|
||||
# No extra dump for the missing-key case since all keys present
|
||||
if hasattr(config, "first_deploy"):
|
||||
assert not config.first_deploy
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=False)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
)
|
||||
def test_first_deploy_flag_set(
|
||||
self, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""first_deploy should be True when config file is created for the first time."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"version": 1},
|
||||
)
|
||||
assert hasattr(config, "first_deploy")
|
||||
assert config.first_deploy is True
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"x": null}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_none_value_replaced_with_default(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""When a config value is null, it is replaced by the default."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"x": "fallback"},
|
||||
)
|
||||
assert config["x"] == "fallback"
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"a": 1, "c": 3}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_missing_keys_inserted_from_default(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""Keys present in default but missing from file are inserted."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"a": 1, "b": 2, "c": 3},
|
||||
)
|
||||
assert config["a"] == 1
|
||||
assert config["b"] == 2 # inserted from default
|
||||
assert config["c"] == 3
|
||||
|
||||
|
||||
class TestAstrBotConfigOperations:
|
||||
"""Dot-notation access, save, and delete."""
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"key": "val"}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_getattr_existing_key(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""__getattr__ returns the value for an existing key."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"key": "val"},
|
||||
)
|
||||
assert config.key == "val"
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"key": "val"}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_getattr_missing_key_returns_none(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""__getattr__ returns None for a missing key."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"key": "val"},
|
||||
)
|
||||
assert config.non_existent is None
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"key": "val"}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_setattr_updates_dict(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""__setattr__ stores the value in the dict."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"key": "val"},
|
||||
)
|
||||
config.new_field = 99
|
||||
assert config["new_field"] == 99
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"temp": "x"}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_delattr_removes_from_dict_and_saves(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""__delattr__ removes the key from the dict and triggers save."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"temp": "x"},
|
||||
)
|
||||
del config.temp
|
||||
assert "temp" not in config
|
||||
mock_dump.assert_called()
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"key": "val"}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_delattr_missing_key_raises(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""Deleting a nonexistent key raises AttributeError."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"key": "val"},
|
||||
)
|
||||
with pytest.raises(AttributeError):
|
||||
del config.non_existent
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"key": "val"}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_save_config_writes_to_file(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""save_config serialises self to the config_path."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"key": "val"},
|
||||
)
|
||||
config.new_thing = "test"
|
||||
config.save_config()
|
||||
# The file should be written; at this point json.dump was already called
|
||||
# during __init__ as well, so at least one call after our modification
|
||||
assert mock_dump.call_count >= 1
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=True)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
read_data='{"key": "val"}',
|
||||
)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.json.dump",
|
||||
return_value=None,
|
||||
)
|
||||
def test_save_config_with_replace(
|
||||
self, mock_dump: MagicMock, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""Passing replace_config updates self and writes."""
|
||||
config = AstrBotConfig(
|
||||
config_path="/fake/cmd_config.json",
|
||||
default_config={"key": "val"},
|
||||
)
|
||||
config.save_config(replace_config={"replacement": True})
|
||||
assert config["replacement"] is True
|
||||
|
||||
|
||||
class TestAstrBotConfigSchema:
|
||||
"""Schema-based construction."""
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=False)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
)
|
||||
def test_schema_object_type(
|
||||
self, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""A schema with type 'object' produces a nested dict."""
|
||||
schema = {
|
||||
"nested": {"type": "object", "items": {"a": {"type": "int"}}},
|
||||
}
|
||||
config = AstrBotConfig(config_path="/fake/cfg.json", schema=schema)
|
||||
assert config.nested == {"a": 0}
|
||||
|
||||
@patch("astrbot.core.config.astrbot_config.os.path.exists", return_value=False)
|
||||
@patch(
|
||||
"astrbot.core.config.astrbot_config.open",
|
||||
new_callable=mock_open,
|
||||
)
|
||||
def test_schema_unsupported_type_raises(
|
||||
self, mock_file: MagicMock, mock_exists: MagicMock
|
||||
):
|
||||
"""An unsupported type in the schema raises TypeError."""
|
||||
schema = {"bad": {"type": "unknown_type"}}
|
||||
with pytest.raises(TypeError, match="不受支持的配置类型"):
|
||||
AstrBotConfig(config_path="/fake/cfg.json", schema=schema)
|
||||
|
||||
|
||||
class TestAstrBotConfigCheckExist:
|
||||
"""check_exist edge cases."""
|
||||
|
||||
def test_check_exist_empty_path(self):
|
||||
"""Return False when config_path is falsy."""
|
||||
config = AstrBotConfig.__new__(AstrBotConfig)
|
||||
object.__setattr__(config, "config_path", "")
|
||||
assert config.check_exist() is False
|
||||
@@ -0,0 +1,265 @@
|
||||
"""Mock-based unit tests for AstrBotConfigManager."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid as uuid_mod
|
||||
from unittest.mock import MagicMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.astrbot_config_mgr import (
|
||||
AstrBotConfigManager,
|
||||
ConfInfo,
|
||||
DEFAULT_CONFIG_CONF_INFO,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_default_config():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_ucr():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_sp():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def acm(mock_default_config, mock_ucr, mock_sp):
|
||||
"""Build an AstrBotConfigManager with all dependencies mocked."""
|
||||
with patch.object(AstrBotConfigManager, "_load_all_configs", return_value=None):
|
||||
acm = AstrBotConfigManager(mock_default_config, mock_ucr, mock_sp)
|
||||
return acm
|
||||
|
||||
|
||||
class TestAstrBotConfigManagerConstruction:
|
||||
"""Construction and initialisation."""
|
||||
|
||||
def test_init_stores_dependencies(self, acm, mock_default_config, mock_ucr, mock_sp):
|
||||
assert acm.sp is mock_sp
|
||||
assert acm.ucr is mock_ucr
|
||||
assert acm.confs["default"] is mock_default_config
|
||||
|
||||
def test_init_sets_abconf_data_to_none(self, acm):
|
||||
assert acm.abconf_data is None
|
||||
|
||||
def test_default_conf_property(self, acm, mock_default_config):
|
||||
assert acm.default_conf is mock_default_config
|
||||
|
||||
|
||||
class TestGetConf:
|
||||
"""get_conf method."""
|
||||
|
||||
def test_get_conf_returns_default_when_umo_none(self, acm):
|
||||
conf = acm.get_conf(None)
|
||||
assert conf is acm.confs["default"]
|
||||
|
||||
def test_get_conf_returns_default_when_umo_not_mapped(self, acm, mock_ucr):
|
||||
mock_ucr.get_conf_id_for_umop.return_value = None
|
||||
conf = acm.get_conf("test:group:123")
|
||||
assert conf is acm.confs["default"]
|
||||
|
||||
def test_get_conf_returns_mapped(self, acm, mock_ucr, mock_sp):
|
||||
conf_id = "uuid-abc"
|
||||
mock_ucr.get_conf_id_for_umop.return_value = conf_id
|
||||
mock_sp.get.return_value = {conf_id: {"path": "abconf_uuid-abc.json", "name": "test"}}
|
||||
mock_conf = MagicMock()
|
||||
acm.confs[conf_id] = mock_conf
|
||||
conf = acm.get_conf("qq:group:456")
|
||||
assert conf is mock_conf
|
||||
|
||||
def test_get_conf_fallback_when_mapped_not_loaded(self, acm, mock_ucr, mock_sp, mock_default_config):
|
||||
conf_id = "uuid-missing"
|
||||
mock_ucr.get_conf_id_for_umop.return_value = conf_id
|
||||
mock_sp.get.return_value = {conf_id: {"path": "nope.json", "name": "x"}}
|
||||
conf = acm.get_conf("qq:group:789")
|
||||
assert conf is mock_default_config
|
||||
|
||||
|
||||
class TestGetConfInfo:
|
||||
"""get_conf_info method."""
|
||||
|
||||
def test_get_conf_info_returns_default_when_unmapped(self, acm, mock_ucr):
|
||||
mock_ucr.get_conf_id_for_umop.return_value = None
|
||||
info = acm.get_conf_info("qq:group:1")
|
||||
assert info["id"] == "default"
|
||||
|
||||
def test_get_conf_info_returns_mapped_meta(self, acm, mock_ucr, mock_sp):
|
||||
conf_id = "uuid-mapped"
|
||||
mock_ucr.get_conf_id_for_umop.return_value = conf_id
|
||||
mock_sp.get.return_value = {conf_id: {"path": "cfg.json", "name": "MyCfg"}}
|
||||
info = acm.get_conf_info("qq:group:2")
|
||||
assert info["id"] == conf_id
|
||||
assert info["path"] == "cfg.json"
|
||||
assert "umop" not in info
|
||||
|
||||
|
||||
class TestGetConfList:
|
||||
"""get_conf_list method."""
|
||||
|
||||
def test_get_conf_list_includes_default(self, acm):
|
||||
acm.abconf_data = {}
|
||||
lst = acm.get_conf_list()
|
||||
assert DEFAULT_CONFIG_CONF_INFO in lst
|
||||
|
||||
def test_get_conf_list_returns_all_abconfs(self, acm, mock_sp):
|
||||
mock_sp.get.return_value = {
|
||||
"u1": {"path": "a.json", "name": "A"},
|
||||
"u2": {"path": "b.json", "name": "B"},
|
||||
}
|
||||
acm.abconf_data = mock_sp.get.return_value
|
||||
lst = acm.get_conf_list()
|
||||
ids = {item["id"] for item in lst}
|
||||
assert "u1" in ids
|
||||
assert "u2" in ids
|
||||
assert "default" in ids
|
||||
|
||||
def test_get_conf_list_skips_non_dict(self, acm, mock_sp):
|
||||
mock_sp.get.return_value = {
|
||||
"u1": {"path": "a.json", "name": "A"},
|
||||
"u2": "not a dict",
|
||||
}
|
||||
acm.abconf_data = mock_sp.get.return_value
|
||||
lst = acm.get_conf_list()
|
||||
assert len(lst) == 2 # only u1 + default
|
||||
|
||||
|
||||
class TestCreateConf:
|
||||
"""create_conf method."""
|
||||
|
||||
@patch("astrbot.core.astrbot_config_mgr.uuid.uuid4", return_value=uuid_mod.UUID("00000000-0000-0000-0000-000000000001"))
|
||||
@patch("astrbot.core.astrbot_config_mgr.AstrBotConfig")
|
||||
@patch("astrbot.core.astrbot_config_mgr.get_astrbot_config_path", return_value="/cfg")
|
||||
def test_create_conf_creates_and_saves(
|
||||
self,
|
||||
mock_get_path,
|
||||
mock_Config,
|
||||
mock_uuid,
|
||||
acm,
|
||||
):
|
||||
mock_conf_instance = MagicMock()
|
||||
mock_Config.return_value = mock_conf_instance
|
||||
conf_id = acm.create_conf(config={"key": "val"}, name="myname")
|
||||
mock_Config.assert_called_once()
|
||||
mock_conf_instance.save_config.assert_called_once()
|
||||
assert conf_id in acm.confs
|
||||
assert acm.confs[conf_id] is mock_conf_instance
|
||||
|
||||
|
||||
class TestDeleteConf:
|
||||
"""delete_conf method."""
|
||||
|
||||
def test_delete_conf_raises_on_default(self, acm):
|
||||
with pytest.raises(ValueError, match="不能删除默认配置文件"):
|
||||
acm.delete_conf("default")
|
||||
|
||||
def test_delete_conf_returns_false_when_not_found(self, acm, mock_sp):
|
||||
mock_sp.get.return_value = {}
|
||||
result = acm.delete_conf("nonexistent")
|
||||
assert result is False
|
||||
|
||||
@patch("astrbot.core.astrbot_config_mgr.os.remove")
|
||||
@patch("astrbot.core.astrbot_config_mgr.os.path.exists", return_value=True)
|
||||
@patch("astrbot.core.astrbot_config_mgr.get_astrbot_config_path", return_value="/cfg")
|
||||
def test_delete_conf_removes_file_and_mapping(
|
||||
self,
|
||||
mock_get_path,
|
||||
mock_exists,
|
||||
mock_remove,
|
||||
acm,
|
||||
mock_sp,
|
||||
):
|
||||
conf_id = "uuid-to-delete"
|
||||
mock_sp.get.return_value = {conf_id: {"path": "abconf_uuid-to-delete.json", "name": "x"}}
|
||||
acm.abconf_data = mock_sp.get.return_value
|
||||
|
||||
result = acm.delete_conf(conf_id)
|
||||
assert result is True
|
||||
mock_remove.assert_called_once()
|
||||
mock_sp.put.assert_called()
|
||||
|
||||
@patch("astrbot.core.astrbot_config_mgr.os.path.exists", return_value=False)
|
||||
@patch("astrbot.core.astrbot_config_mgr.get_astrbot_config_path", return_value="/cfg")
|
||||
def test_delete_conf_handles_missing_file(
|
||||
self, mock_get_path, mock_exists, acm, mock_sp
|
||||
):
|
||||
conf_id = "uuid-missing-file"
|
||||
mock_sp.get.return_value = {conf_id: {"path": "gone.json", "name": "x"}}
|
||||
acm.abconf_data = mock_sp.get.return_value
|
||||
acm.confs[conf_id] = MagicMock()
|
||||
|
||||
result = acm.delete_conf(conf_id)
|
||||
assert result is True
|
||||
assert conf_id not in acm.confs
|
||||
|
||||
|
||||
class TestUpdateConfInfo:
|
||||
"""update_conf_info method."""
|
||||
|
||||
def test_update_raises_on_default(self, acm):
|
||||
with pytest.raises(ValueError, match="不能更新"):
|
||||
acm.update_conf_info("default", name="new")
|
||||
|
||||
def test_update_returns_false_when_not_found(self, acm, mock_sp):
|
||||
mock_sp.get.return_value = {}
|
||||
result = acm.update_conf_info("nonexistent", name="new")
|
||||
assert result is False
|
||||
|
||||
def test_update_renames(self, acm, mock_sp):
|
||||
conf_id = "uuid-rename"
|
||||
mock_sp.get.return_value = {conf_id: {"path": "x.json", "name": "old"}}
|
||||
acm.abconf_data = mock_sp.get.return_value
|
||||
result = acm.update_conf_info(conf_id, name="new_name")
|
||||
assert result is True
|
||||
assert acm.abconf_data[conf_id]["name"] == "new_name"
|
||||
|
||||
|
||||
class TestG:
|
||||
"""g (generic getter) method."""
|
||||
|
||||
def test_g_without_umo_uses_default(self, acm, mock_default_config):
|
||||
mock_default_config.get.return_value = "fallback"
|
||||
val = acm.g(umo=None, key="missing")
|
||||
mock_default_config.get.assert_called_with("missing", None)
|
||||
assert val == "fallback"
|
||||
|
||||
def test_g_with_umo_uses_get_conf(self, acm):
|
||||
fake_conf = MagicMock()
|
||||
fake_conf.get.return_value = 42
|
||||
acm.get_conf = MagicMock(return_value=fake_conf)
|
||||
val = acm.g(umo="qq:group:1", key="some.setting")
|
||||
assert val == 42
|
||||
fake_conf.get.assert_called_with("some.setting", None)
|
||||
|
||||
|
||||
class TestSaveConfMapping:
|
||||
"""_save_conf_mapping internal method."""
|
||||
|
||||
def test_save_conf_mapping_stores_and_updates_abconf_data(self, acm, mock_sp):
|
||||
mock_sp.get.return_value = {}
|
||||
acm._save_conf_mapping(abconf_path="new.json", abconf_id="new-id", abconf_name="display")
|
||||
mock_sp.put.assert_called()
|
||||
assert "new-id" in acm.abconf_data
|
||||
|
||||
|
||||
class TestLoadConfMappingEdgeCases:
|
||||
"""_load_conf_mapping edge cases."""
|
||||
|
||||
def test_load_conf_mapping_with_invalid_umo_str(self, acm, mock_ucr):
|
||||
"""An invalid umo string that can't be parsed as MessageSession returns default."""
|
||||
mock_ucr.get_conf_id_for_umop.side_effect = Exception("parse error")
|
||||
info = acm._load_conf_mapping("bad_format")
|
||||
assert info["id"] == "default"
|
||||
|
||||
def test_load_conf_mapping_checks_meta_is_dict(self, acm, mock_ucr, mock_sp):
|
||||
"""If abconf metadata is not a dict, returns default."""
|
||||
conf_id = "uuid-non-dict"
|
||||
mock_ucr.get_conf_id_for_umop.return_value = conf_id
|
||||
mock_sp.get.return_value = {conf_id: "not a dict"}
|
||||
info = acm._load_conf_mapping("qq:friend:1")
|
||||
assert info["id"] == "default"
|
||||
@@ -0,0 +1,267 @@
|
||||
"""Mock-based unit tests for AstrBot builtin star (Main)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.api.event import AstrMessageEvent
|
||||
from astrbot.api.message_components import Image, Plain
|
||||
from astrbot.api.provider import LLMResponse, ProviderRequest
|
||||
from astrbot.builtin_stars.astrbot.main import Main
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_star_context():
|
||||
ctx = MagicMock()
|
||||
ctx.astrbot_config_mgr = MagicMock()
|
||||
return ctx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def star(mock_star_context):
|
||||
return Main(mock_star_context)
|
||||
|
||||
|
||||
class TestMainConstruction:
|
||||
"""Construction and LTM initialisation."""
|
||||
|
||||
def test_init_stores_context(self, star, mock_star_context):
|
||||
assert star.context is mock_star_context
|
||||
|
||||
@patch("astrbot.builtin_stars.astrbot.main.LongTermMemory")
|
||||
def test_init_creates_ltm(self, mock_ltm_cls, mock_star_context):
|
||||
mock_ltm_instance = MagicMock()
|
||||
mock_ltm_cls.return_value = mock_ltm_instance
|
||||
s = Main(mock_star_context)
|
||||
mock_ltm_cls.assert_called_once_with(
|
||||
mock_star_context.astrbot_config_mgr, mock_star_context
|
||||
)
|
||||
assert s.ltm is mock_ltm_instance
|
||||
|
||||
@patch("astrbot.builtin_stars.astrbot.main.logger")
|
||||
@patch("astrbot.builtin_stars.astrbot.main.LongTermMemory", side_effect=ValueError("fail"))
|
||||
def test_init_handles_ltm_failure(self, mock_ltm_cls, mock_logger, mock_star_context):
|
||||
s = Main(mock_star_context)
|
||||
assert s.ltm is None
|
||||
mock_logger.error.assert_called()
|
||||
|
||||
|
||||
class TestLtmEnabled:
|
||||
"""ltm_enabled helper."""
|
||||
|
||||
def test_ltm_enabled_true_when_group_icl(self, star):
|
||||
event = MagicMock()
|
||||
event.unified_msg_origin = "qq:group:1"
|
||||
star.context.get_config.return_value = {
|
||||
"provider_ltm_settings": {
|
||||
"group_icl_enable": True,
|
||||
"active_reply": {"enable": False},
|
||||
},
|
||||
}
|
||||
assert star.ltm_enabled(event) is True
|
||||
|
||||
def test_ltm_enabled_true_when_active_reply(self, star):
|
||||
event = MagicMock()
|
||||
event.unified_msg_origin = "qq:group:1"
|
||||
star.context.get_config.return_value = {
|
||||
"provider_ltm_settings": {
|
||||
"group_icl_enable": False,
|
||||
"active_reply": {"enable": True},
|
||||
},
|
||||
}
|
||||
assert star.ltm_enabled(event) is True
|
||||
|
||||
def test_ltm_enabled_false_when_both_disabled(self, star):
|
||||
event = MagicMock()
|
||||
event.unified_msg_origin = "qq:group:1"
|
||||
star.context.get_config.return_value = {
|
||||
"provider_ltm_settings": {
|
||||
"group_icl_enable": False,
|
||||
"active_reply": {"enable": False},
|
||||
},
|
||||
}
|
||||
assert star.ltm_enabled(event) is False
|
||||
|
||||
|
||||
class TestOnMessage:
|
||||
"""on_message handler."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_message_skips_when_no_plain_or_image(self, star):
|
||||
event = MagicMock()
|
||||
event.message_obj.message = [MagicMock(spec=object)] # neither Plain nor Image
|
||||
gen = star.on_message(event)
|
||||
items = [item async for item in gen]
|
||||
assert items == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_message_skips_when_ltm_disabled(self, star):
|
||||
event = MagicMock()
|
||||
event.message_obj.message = [Plain(text="hello")]
|
||||
star.ltm_enabled = MagicMock(return_value=False)
|
||||
gen = star.on_message(event)
|
||||
items = [item async for item in gen]
|
||||
assert items == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.object(Main, "ltm_enabled", return_value=True)
|
||||
async def test_on_message_records_context(self, mock_enabled, star):
|
||||
event = MagicMock()
|
||||
event.message_obj.message = [Plain(text="hello")]
|
||||
star.ltm = MagicMock()
|
||||
star.ltm.need_active_reply = AsyncMock(return_value=False)
|
||||
star.ltm.handle_message = AsyncMock()
|
||||
|
||||
star.context.get_config.return_value = {
|
||||
"provider_ltm_settings": {"group_icl_enable": True, "active_reply": {"enable": False}},
|
||||
}
|
||||
gen = star.on_message(event)
|
||||
items = [item async for item in gen]
|
||||
star.ltm.handle_message.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.object(Main, "ltm_enabled", return_value=True)
|
||||
async def test_on_message_active_reply_triggers_llm(self, mock_enabled, star):
|
||||
event = MagicMock()
|
||||
event.message_obj.message = [Plain(text="hello")]
|
||||
event.message_str = "hello"
|
||||
event.session_id = "sess-1"
|
||||
event.unified_msg_origin = "qq:group:1"
|
||||
|
||||
star.ltm = MagicMock()
|
||||
star.ltm.need_active_reply = AsyncMock(return_value=True)
|
||||
star.ltm.handle_message = AsyncMock()
|
||||
|
||||
awaitables = []
|
||||
|
||||
def request_llm(prompt, session_id, conversation):
|
||||
nonlocal awaitables
|
||||
awaitables.append((prompt, session_id, conversation))
|
||||
return MagicMock()
|
||||
|
||||
event.request_llm = MagicMock(side_effect=request_llm)
|
||||
|
||||
provider = MagicMock()
|
||||
provider.text_chat = AsyncMock()
|
||||
star.context.get_using_provider.return_value = provider
|
||||
star.context.conversation_manager.get_curr_conversation_id = AsyncMock(return_value="cid-1")
|
||||
conv = MagicMock()
|
||||
star.context.conversation_manager.get_conversation = AsyncMock(return_value=conv)
|
||||
|
||||
star.context.get_config.return_value = {
|
||||
"provider_ltm_settings": {"group_icl_enable": True, "active_reply": {"enable": True}},
|
||||
}
|
||||
|
||||
gen = star.on_message(event)
|
||||
items = [item async for item in gen]
|
||||
assert len(awaitables) == 1
|
||||
assert awaitables[0][0] == "hello"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch.object(Main, "ltm_enabled", return_value=True)
|
||||
async def test_on_message_logs_when_no_provider(self, mock_enabled, star):
|
||||
event = MagicMock()
|
||||
event.message_obj.message = [Plain(text="hi")]
|
||||
star.ltm = MagicMock()
|
||||
star.ltm.need_active_reply = AsyncMock(return_value=True)
|
||||
star.ltm.handle_message = AsyncMock()
|
||||
star.context.get_using_provider.return_value = None
|
||||
star.context.get_config.return_value = {
|
||||
"provider_ltm_settings": {"group_icl_enable": True, "active_reply": {"enable": True}},
|
||||
}
|
||||
gen = star.on_message(event)
|
||||
items = [item async for item in gen]
|
||||
assert items == []
|
||||
|
||||
|
||||
class TestDecorateLlmReq:
|
||||
"""Decorating LLM requests."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decorate_llm_req_calls_ltm(self, star):
|
||||
event = MagicMock()
|
||||
req = ProviderRequest(prompt="test")
|
||||
star.ltm = MagicMock()
|
||||
star.ltm.on_req_llm = AsyncMock()
|
||||
star.ltm_enabled = MagicMock(return_value=True)
|
||||
await star.decorate_llm_req(event, req)
|
||||
star.ltm.on_req_llm.assert_awaited_once_with(event, req)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decorate_llm_req_skips_when_disabled(self, star):
|
||||
event = MagicMock()
|
||||
req = MagicMock()
|
||||
star.ltm = MagicMock()
|
||||
star.ltm_enabled = MagicMock(return_value=False)
|
||||
await star.decorate_llm_req(event, req)
|
||||
star.ltm.on_req_llm.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decorate_llm_req_skips_when_ltm_none(self, star):
|
||||
star.ltm = None
|
||||
event = MagicMock()
|
||||
req = MagicMock()
|
||||
await star.decorate_llm_req(event, req)
|
||||
|
||||
|
||||
class TestRecordLlmResp:
|
||||
"""Recording LLM responses."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_record_llm_resp_calls_ltm(self, star):
|
||||
event = MagicMock()
|
||||
resp = MagicMock(spec=LLMResponse)
|
||||
star.ltm = MagicMock()
|
||||
star.ltm.after_req_llm = AsyncMock()
|
||||
star.ltm_enabled = MagicMock(return_value=True)
|
||||
await star.record_llm_resp_to_ltm(event, resp)
|
||||
star.ltm.after_req_llm.assert_awaited_once_with(event, resp)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_record_llm_resp_skips_when_disabled(self, star):
|
||||
event = MagicMock()
|
||||
resp = MagicMock(spec=LLMResponse)
|
||||
star.ltm = MagicMock()
|
||||
star.ltm_enabled = MagicMock(return_value=False)
|
||||
await star.record_llm_resp_to_ltm(event, resp)
|
||||
star.ltm.after_req_llm.assert_not_awaited()
|
||||
|
||||
|
||||
class TestAfterMessageSent:
|
||||
"""After-message-sent handler."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_after_message_sent_cleans_session(self, star):
|
||||
event = MagicMock()
|
||||
event.get_extra.return_value = True
|
||||
star.ltm = MagicMock()
|
||||
star.ltm.remove_session = AsyncMock(return_value=3)
|
||||
star.ltm_enabled = MagicMock(return_value=True)
|
||||
await star.after_message_sent(event)
|
||||
event.get_extra.assert_called_with("_clean_ltm_session", False)
|
||||
star.ltm.remove_session.assert_awaited_once_with(event)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_after_message_sent_skips_when_not_flagged(self, star):
|
||||
event = MagicMock()
|
||||
event.get_extra.return_value = False
|
||||
star.ltm = MagicMock()
|
||||
star.ltm_enabled = MagicMock(return_value=True)
|
||||
await star.after_message_sent(event)
|
||||
star.ltm.remove_session.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_after_message_sent_skips_when_ltm_disabled(self, star):
|
||||
event = MagicMock()
|
||||
star.ltm = MagicMock()
|
||||
star.ltm_enabled = MagicMock(return_value=False)
|
||||
await star.after_message_sent(event)
|
||||
star.ltm.remove_session.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_after_message_sent_skips_when_ltm_none(self, star):
|
||||
star.ltm = None
|
||||
event = MagicMock()
|
||||
await star.after_message_sent(event)
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Unit tests for astrbot.core.utils.command_parser.
|
||||
|
||||
Covers CommandTokens data-holder and CommandParserMixin parse/regex helpers.
|
||||
"""
|
||||
|
||||
from astrbot.core.utils.command_parser import CommandParserMixin, CommandTokens
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# CommandTokens
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCommandTokens:
|
||||
def test_default_attributes(self):
|
||||
tokens = CommandTokens()
|
||||
assert tokens.tokens == []
|
||||
assert tokens.len == 0
|
||||
|
||||
def test_get_returns_none_for_out_of_bounds_negative(self):
|
||||
tokens = CommandTokens()
|
||||
assert tokens.get(-1) is None
|
||||
|
||||
def test_get_returns_none_for_out_of_bounds_positive(self):
|
||||
tokens = CommandTokens()
|
||||
tokens.tokens = ["a", "b"]
|
||||
tokens.len = 2
|
||||
assert tokens.get(2) is None
|
||||
assert tokens.get(100) is None
|
||||
|
||||
def test_get_returns_stripped_token(self):
|
||||
tokens = CommandTokens()
|
||||
tokens.tokens = [" hello ", "world "]
|
||||
tokens.len = 2
|
||||
assert tokens.get(0) == "hello"
|
||||
assert tokens.get(1) == "world"
|
||||
|
||||
def test_get_from_empty_list(self):
|
||||
tokens = CommandTokens()
|
||||
tokens.tokens = []
|
||||
tokens.len = 0
|
||||
assert tokens.get(0) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# CommandParserMixin
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _ConcreteParser(CommandParserMixin):
|
||||
"""Minimal concrete subclass so the mixin can be instantiated."""
|
||||
|
||||
|
||||
class TestCommandParserMixin:
|
||||
def setup_method(self) -> None:
|
||||
self.parser = _ConcreteParser()
|
||||
|
||||
def test_parse_commands_returns_command_tokens(self):
|
||||
result = self.parser.parse_commands("hello world")
|
||||
assert isinstance(result, CommandTokens)
|
||||
assert result.len == 2
|
||||
|
||||
def test_parse_commands_splits_on_whitespace(self):
|
||||
result = self.parser.parse_commands("one two\tthree\nfour")
|
||||
assert result.tokens == ["one", "two", "three", "four"]
|
||||
assert result.len == 4
|
||||
|
||||
def test_parse_commands_empty_string_yields_single_empty_token(self):
|
||||
result = self.parser.parse_commands("")
|
||||
assert result.tokens == [""]
|
||||
assert result.len == 1
|
||||
|
||||
def test_parse_commands_only_whitespace(self):
|
||||
result = self.parser.parse_commands(" \t \n ")
|
||||
assert result.tokens == ["", ""] # \S+ splits on whitespace runs
|
||||
# Actually re.split(r"\s+", " \t \n ") = ['', '']
|
||||
|
||||
def test_regex_match_returns_true_on_match(self):
|
||||
assert self.parser.regex_match("hello world", "hello") is True
|
||||
|
||||
def test_regex_match_returns_false_on_no_match(self):
|
||||
assert self.parser.regex_match("hello world", "goodbye") is False
|
||||
|
||||
def test_regex_match_multiline(self):
|
||||
text = "line1\nline2\nline3"
|
||||
assert self.parser.regex_match(text, "^line2$") is True
|
||||
|
||||
def test_regex_match_with_digits(self):
|
||||
assert self.parser.regex_match("abc123def", r"\d+") is True
|
||||
|
||||
def test_regex_match_with_special_regex_chars(self):
|
||||
"""Special regex characters should be treated as regex, not literal."""
|
||||
assert self.parser.regex_match("price is 10.50", r"\d+\.\d+") is True
|
||||
|
||||
def test_regex_match_empty_pattern(self):
|
||||
"""An empty pattern matches any string (re.search always finds '')."""
|
||||
assert self.parser.regex_match("anything", "") is True
|
||||
@@ -0,0 +1,529 @@
|
||||
"""Unit tests for astrbot.core.star.context.
|
||||
|
||||
Tests Context methods with mock-based isolation.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, PropertyMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.star.context import Context
|
||||
from astrbot.core.star.star import StarMetadata, star_map, star_registry
|
||||
from astrbot.core.star.star_handler import (
|
||||
EventType,
|
||||
StarHandlerMetadata,
|
||||
star_handlers_registry,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_dependencies():
|
||||
"""Create fully isolated mocks for all Context constructor dependencies."""
|
||||
return {
|
||||
"event_queue": MagicMock(),
|
||||
"config": MagicMock(),
|
||||
"db": MagicMock(),
|
||||
"provider_manager": MagicMock(),
|
||||
"platform_manager": MagicMock(),
|
||||
"conversation_manager": MagicMock(),
|
||||
"message_history_manager": MagicMock(),
|
||||
"persona_manager": MagicMock(),
|
||||
"astrbot_config_mgr": MagicMock(),
|
||||
"knowledge_base_manager": MagicMock(),
|
||||
"cron_manager": MagicMock(),
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def context(mock_dependencies):
|
||||
"""Create a Context with all mocks."""
|
||||
return Context(
|
||||
mock_dependencies["event_queue"],
|
||||
mock_dependencies["config"],
|
||||
mock_dependencies["db"],
|
||||
mock_dependencies["provider_manager"],
|
||||
mock_dependencies["platform_manager"],
|
||||
mock_dependencies["conversation_manager"],
|
||||
mock_dependencies["message_history_manager"],
|
||||
mock_dependencies["persona_manager"],
|
||||
mock_dependencies["astrbot_config_mgr"],
|
||||
mock_dependencies["knowledge_base_manager"],
|
||||
mock_dependencies["cron_manager"],
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constructor / Init
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestContextInit:
|
||||
"""Context constructor and attribute initialization."""
|
||||
|
||||
def test_init_sets_attributes(self, context, mock_dependencies):
|
||||
"""All constructor arguments are stored as instance attributes."""
|
||||
assert context._event_queue is mock_dependencies["event_queue"]
|
||||
assert context._config is mock_dependencies["config"]
|
||||
assert context._db is mock_dependencies["db"]
|
||||
assert context.provider_manager is mock_dependencies["provider_manager"]
|
||||
assert context.platform_manager is mock_dependencies["platform_manager"]
|
||||
assert context.conversation_manager is mock_dependencies["conversation_manager"]
|
||||
assert context.message_history_manager is mock_dependencies["message_history_manager"]
|
||||
assert context.persona_manager is mock_dependencies["persona_manager"]
|
||||
assert context.astrbot_config_mgr is mock_dependencies["astrbot_config_mgr"]
|
||||
assert context.kb_manager is mock_dependencies["knowledge_base_manager"]
|
||||
assert context.cron_manager is mock_dependencies["cron_manager"]
|
||||
|
||||
def test_init_sets_empty_registrations(self, context):
|
||||
"""Runtime registration containers start empty."""
|
||||
assert context._registered_web_apis == []
|
||||
assert context._register_tasks == []
|
||||
assert context._star_manager is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_using_provider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetUsingProvider:
|
||||
"""Context.get_using_provider() behavior."""
|
||||
|
||||
def test_returns_provider_when_found(self, context, mock_dependencies):
|
||||
"""get_using_provider returns the provider from provider_manager."""
|
||||
mock_provider = MagicMock()
|
||||
mock_provider.__class__.__name__ = "Provider"
|
||||
mock_dependencies["provider_manager"].get_using_provider.return_value = (
|
||||
mock_provider
|
||||
)
|
||||
result = context.get_using_provider("test_umo")
|
||||
assert result is mock_provider
|
||||
mock_dependencies["provider_manager"].get_using_provider.assert_called_once()
|
||||
|
||||
def test_returns_none_when_not_found(self, context, mock_dependencies):
|
||||
"""get_using_provider returns None when no provider is available."""
|
||||
mock_dependencies["provider_manager"].get_using_provider.return_value = None
|
||||
result = context.get_using_provider("test_umo")
|
||||
assert result is None
|
||||
|
||||
def test_raises_value_error_on_wrong_type(self, context, mock_dependencies):
|
||||
"""get_using_provider raises ValueError when provider is wrong type."""
|
||||
mock_dependencies["provider_manager"].get_using_provider.return_value = (
|
||||
"not_a_provider"
|
||||
)
|
||||
with pytest.raises(ValueError, match="类型不正确"):
|
||||
context.get_using_provider("test_umo")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_config
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetConfig:
|
||||
"""Context.get_config() behavior."""
|
||||
|
||||
def test_returns_default_config_without_umo(self, context, mock_dependencies):
|
||||
"""get_config() returns _config when umo is None."""
|
||||
result = context.get_config()
|
||||
assert result is mock_dependencies["config"]
|
||||
|
||||
def test_returns_umo_config_when_umo_provided(self, context, mock_dependencies):
|
||||
"""get_config(umo) delegates to astrbot_config_mgr."""
|
||||
mock_umo_config = MagicMock()
|
||||
mock_dependencies["astrbot_config_mgr"].get_conf.return_value = mock_umo_config
|
||||
result = context.get_config("test_umo")
|
||||
assert result is mock_umo_config
|
||||
mock_dependencies["astrbot_config_mgr"].get_conf.assert_called_once_with(
|
||||
"test_umo"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_registered_star / get_all_stars
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetRegisteredStar:
|
||||
"""Context.get_registered_star() behavior."""
|
||||
|
||||
@patch("astrbot.core.star.context.star_registry", new_callable=list)
|
||||
def test_finds_star_by_name(self, mock_registry, context):
|
||||
"""get_registered_star returns the matching StarMetadata."""
|
||||
s1 = MagicMock(spec=StarMetadata, name="plugin_a")
|
||||
s1.name = "plugin_a"
|
||||
s2 = MagicMock(spec=StarMetadata, name="plugin_b")
|
||||
s2.name = "plugin_b"
|
||||
mock_registry.extend([s1, s2])
|
||||
result = context.get_registered_star("plugin_a")
|
||||
assert result is s1
|
||||
|
||||
@patch("astrbot.core.star.context.star_registry", new_callable=list)
|
||||
def test_returns_none_when_not_found(self, mock_registry, context):
|
||||
"""get_registered_star returns None when no plugin matches."""
|
||||
mock_registry.clear()
|
||||
result = context.get_registered_star("nonexistent")
|
||||
assert result is None
|
||||
|
||||
def test_get_all_stars_returns_registry(self, context):
|
||||
"""get_all_stars returns the module-level star_registry list."""
|
||||
assert context.get_all_stars() is star_registry
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# register_commands
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRegisterCommands:
|
||||
"""Context.register_commands() behavior."""
|
||||
|
||||
def test_registers_command_handler(self, context):
|
||||
"""register_commands creates a StarHandlerMetadata and appends it."""
|
||||
async def fake_handler():
|
||||
pass
|
||||
|
||||
fake_handler.__module__ = "data.plugins.test.main"
|
||||
fake_handler.__qualname__ = "my_command"
|
||||
context.register_commands(
|
||||
star_name="test_star",
|
||||
command_name="/hello",
|
||||
desc="Says hello",
|
||||
priority=5,
|
||||
awaitable=fake_handler,
|
||||
)
|
||||
|
||||
handlers = star_handlers_registry.get_handlers_by_module_name(
|
||||
"data.plugins.test.main"
|
||||
)
|
||||
assert len(handlers) == 1
|
||||
md = handlers[0]
|
||||
assert md.event_type == EventType.AdapterMessageEvent
|
||||
assert md.desc == "Says hello"
|
||||
assert md.handler is fake_handler
|
||||
|
||||
def test_registers_regex_command(self, context):
|
||||
"""register_commands with use_regex=True adds a RegexFilter."""
|
||||
async def fake_handler():
|
||||
pass
|
||||
|
||||
fake_handler.__module__ = "data.plugins.test.main"
|
||||
fake_handler.__qualname__ = "regex_cmd"
|
||||
context.register_commands(
|
||||
star_name="test_star",
|
||||
command_name=r"hello.*",
|
||||
desc="Regex command",
|
||||
priority=1,
|
||||
awaitable=fake_handler,
|
||||
use_regex=True,
|
||||
)
|
||||
|
||||
handlers = star_handlers_registry.get_handlers_by_module_name(
|
||||
"data.plugins.test.main"
|
||||
)
|
||||
assert len(handlers) == 1
|
||||
md = handlers[0]
|
||||
# Should have a RegexFilter (not a CommandFilter)
|
||||
from astrbot.core.star.filter.regex import RegexFilter
|
||||
|
||||
assert any(isinstance(f, RegexFilter) for f in md.event_filters)
|
||||
|
||||
def test_registers_command_with_ignore_prefix(self, context):
|
||||
"""register_commands with use_regex=False adds a CommandFilter."""
|
||||
async def fake_handler():
|
||||
pass
|
||||
|
||||
fake_handler.__module__ = "data.plugins.test.main"
|
||||
fake_handler.__qualname__ = "cmd_no_regex"
|
||||
context.register_commands(
|
||||
star_name="test_star",
|
||||
command_name="/test",
|
||||
desc="A command",
|
||||
priority=1,
|
||||
awaitable=fake_handler,
|
||||
use_regex=False,
|
||||
)
|
||||
|
||||
handlers = star_handlers_registry.get_handlers_by_module_name(
|
||||
"data.plugins.test.main"
|
||||
)
|
||||
assert len(handlers) == 1
|
||||
md = handlers[0]
|
||||
from astrbot.core.star.filter.command import CommandFilter
|
||||
|
||||
assert any(isinstance(f, CommandFilter) for f in md.event_filters)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# register_web_api
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRegisterWebApi:
|
||||
"""Context.register_web_api() behavior."""
|
||||
|
||||
def test_registers_new_api(self, context):
|
||||
"""register_web_api appends a new web API route."""
|
||||
async def handler():
|
||||
pass
|
||||
|
||||
context.register_web_api(
|
||||
route="/api/test",
|
||||
view_handler=handler,
|
||||
methods=["GET"],
|
||||
desc="Test endpoint",
|
||||
)
|
||||
assert len(context._registered_web_apis) == 1
|
||||
assert context._registered_web_apis[0] == (
|
||||
"/api/test",
|
||||
handler,
|
||||
["GET"],
|
||||
"Test endpoint",
|
||||
)
|
||||
|
||||
def test_replaces_existing_route_with_same_methods(self, context):
|
||||
"""register_web_api replaces a previously registered API with same route and methods."""
|
||||
async def old_handler():
|
||||
pass
|
||||
|
||||
async def new_handler():
|
||||
pass
|
||||
|
||||
context.register_web_api(
|
||||
route="/api/test",
|
||||
view_handler=old_handler,
|
||||
methods=["GET"],
|
||||
desc="Old handler",
|
||||
)
|
||||
context.register_web_api(
|
||||
route="/api/test",
|
||||
view_handler=new_handler,
|
||||
methods=["GET"],
|
||||
desc="New handler",
|
||||
)
|
||||
assert len(context._registered_web_apis) == 1
|
||||
assert context._registered_web_apis[0][1] is new_handler
|
||||
assert context._registered_web_apis[0][3] == "New handler"
|
||||
|
||||
def test_allows_different_methods_on_same_route(self, context):
|
||||
"""register_web_api treats different HTTP methods as separate entries."""
|
||||
async def handler():
|
||||
pass
|
||||
|
||||
context.register_web_api(
|
||||
route="/api/test",
|
||||
view_handler=handler,
|
||||
methods=["GET"],
|
||||
desc="GET handler",
|
||||
)
|
||||
context.register_web_api(
|
||||
route="/api/test",
|
||||
view_handler=handler,
|
||||
methods=["POST"],
|
||||
desc="POST handler",
|
||||
)
|
||||
assert len(context._registered_web_apis) == 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# add_llm_tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestAddLLMTools:
|
||||
"""Context.add_llm_tools() behavior."""
|
||||
|
||||
def test_adds_tool_with_module_path(self, context, mock_dependencies):
|
||||
"""add_llm_tools appends tools and sets handler_module_path."""
|
||||
tool = MagicMock()
|
||||
tool.name = "test_tool"
|
||||
tool.__module__ = "astrbot.builtin_stars.my_plugin.main"
|
||||
|
||||
mock_dependencies["provider_manager"].llm_tools.func_list = []
|
||||
context.add_llm_tools(tool)
|
||||
|
||||
assert tool.handler_module_path == "astrbot.builtin_stars.my_plugin.main"
|
||||
assert tool in mock_dependencies["provider_manager"].llm_tools.func_list
|
||||
|
||||
def test_replaces_existing_tool_with_same_name(self, context, mock_dependencies):
|
||||
"""add_llm_tools replaces an existing tool with the same name."""
|
||||
old_tool = MagicMock()
|
||||
old_tool.name = "dup_tool"
|
||||
old_tool.__module__ = "old.module"
|
||||
|
||||
new_tool = MagicMock()
|
||||
new_tool.name = "dup_tool"
|
||||
new_tool.__module__ = "new.module"
|
||||
|
||||
mock_dependencies["provider_manager"].llm_tools.func_list = [old_tool]
|
||||
context.add_llm_tools(new_tool)
|
||||
|
||||
assert old_tool not in mock_dependencies["provider_manager"].llm_tools.func_list
|
||||
assert new_tool in mock_dependencies["provider_manager"].llm_tools.func_list
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# register_llm_tool (deprecated)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRegisterLLMTool:
|
||||
"""Context.register_llm_tool() (deprecated) behavior."""
|
||||
|
||||
def test_registers_handler_and_adds_func(self, context, mock_dependencies):
|
||||
"""register_llm_tool creates a StarHandlerMetadata and adds to func_list."""
|
||||
async def handler():
|
||||
pass
|
||||
|
||||
handler.__module__ = "data.plugins.p.main"
|
||||
handler.__qualname__ = "my_func"
|
||||
mock_dependencies["provider_manager"].llt = MagicMock()
|
||||
mock_dependencies["provider_manager"].llm_tools.func_list = []
|
||||
|
||||
context.register_llm_tool(
|
||||
name="my_tool",
|
||||
func_args=[{"type": "string", "name": "arg1"}],
|
||||
desc="My tool",
|
||||
func_obj=handler,
|
||||
)
|
||||
|
||||
# Check StarHandlerMetadata was added
|
||||
handlers = star_handlers_registry.get_handlers_by_module_name(
|
||||
"data.plugins.p.main"
|
||||
)
|
||||
assert len(handlers) == 1
|
||||
assert handlers[0].event_type == EventType.OnLLMRequestEvent
|
||||
|
||||
def test_calls_add_func_on_manager(self, context, mock_dependencies):
|
||||
"""register_llm_tool delegates to llm_tools.add_func."""
|
||||
async def handler():
|
||||
pass
|
||||
|
||||
handler.__module__ = "mod"
|
||||
handler.__qualname__ = "fn"
|
||||
mock_dependencies["provider_manager"].llm_tools.func_list = []
|
||||
|
||||
context.register_llm_tool(
|
||||
name="my_tool",
|
||||
func_args=[{"type": "string", "name": "arg1"}],
|
||||
desc="desc",
|
||||
func_obj=handler,
|
||||
)
|
||||
|
||||
mock_dependencies["provider_manager"].llm_tools.add_func.assert_called_once_with(
|
||||
"my_tool",
|
||||
[{"type": "string", "name": "arg1"}],
|
||||
"desc",
|
||||
handler,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# register_task / reset_runtime_registrations
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRegisterTask:
|
||||
"""Context.register_task() behavior."""
|
||||
|
||||
def test_appends_task(self, context):
|
||||
"""register_task appends to _register_tasks."""
|
||||
async def task():
|
||||
pass
|
||||
|
||||
context.register_task(task, "A background task")
|
||||
assert task in context._register_tasks
|
||||
|
||||
|
||||
class TestResetRuntimeRegistrations:
|
||||
"""Context.reset_runtime_registrations() behavior."""
|
||||
|
||||
def test_clears_web_apis_and_tasks(self, context):
|
||||
"""reset_runtime_registrations clears both containers."""
|
||||
async def handler():
|
||||
pass
|
||||
|
||||
async def task():
|
||||
pass
|
||||
|
||||
context.register_web_api("/api/test", handler, ["GET"], "desc")
|
||||
context.register_task(task, "desc")
|
||||
assert len(context._registered_web_apis) == 1
|
||||
assert len(context._register_tasks) == 1
|
||||
|
||||
context.reset_runtime_registrations()
|
||||
|
||||
assert context._registered_web_apis == []
|
||||
assert context._register_tasks == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_platform / get_platform_inst (deprecated)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetPlatform:
|
||||
"""Context.get_platform() (deprecated) and get_platform_inst()."""
|
||||
|
||||
def test_get_platform_by_string_name(self, context, mock_dependencies):
|
||||
"""get_platform finds a platform by its name string."""
|
||||
platform = MagicMock()
|
||||
platform.meta().name = "telegram"
|
||||
mock_dependencies["platform_manager"].platform_insts = [platform]
|
||||
|
||||
result = context.get_platform("telegram")
|
||||
assert result is platform
|
||||
|
||||
def test_get_platform_returns_none_when_not_found(self, context, mock_dependencies):
|
||||
"""get_platform returns None when no platform matches."""
|
||||
mock_dependencies["platform_manager"].platform_insts = []
|
||||
result = context.get_platform("telegram")
|
||||
assert result is None
|
||||
|
||||
def test_get_platform_inst_by_id(self, context, mock_dependencies):
|
||||
"""get_platform_inst finds a platform by its meta id."""
|
||||
platform = MagicMock()
|
||||
platform.meta().id = "test_id"
|
||||
mock_dependencies["platform_manager"].platform_insts = [platform]
|
||||
|
||||
result = context.get_platform_inst("test_id")
|
||||
assert result is platform
|
||||
|
||||
def test_get_platform_inst_returns_none(self, context, mock_dependencies):
|
||||
"""get_platform_inst returns None when no platform matches."""
|
||||
platform = MagicMock()
|
||||
platform.meta().id = "other_id"
|
||||
mock_dependencies["platform_manager"].platform_insts = [platform]
|
||||
result = context.get_platform_inst("nonexistent")
|
||||
assert result is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_db
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetDB:
|
||||
"""Context.get_db() behavior."""
|
||||
|
||||
def test_returns_db(self, context, mock_dependencies):
|
||||
"""get_db returns the _db instance."""
|
||||
assert context.get_db() is mock_dependencies["db"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_event_queue
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetEventQueue:
|
||||
"""Context.get_event_queue() behavior."""
|
||||
|
||||
def test_returns_event_queue(self, context, mock_dependencies):
|
||||
"""get_event_queue returns the _event_queue."""
|
||||
assert context.get_event_queue() is mock_dependencies["event_queue"]
|
||||
@@ -0,0 +1,785 @@
|
||||
"""Unit tests for astrbot.core.conversation_mgr.ConversationManager."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.conversation_mgr import ConversationManager
|
||||
from astrbot.core.db.po import Conversation, ConversationV2
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_conv_v2(
|
||||
conversation_id="conv-1",
|
||||
platform_id="test",
|
||||
user_id="test_user",
|
||||
title="Test Title",
|
||||
persona_id=None,
|
||||
content=None,
|
||||
token_usage=None,
|
||||
):
|
||||
"""Factory for ConversationV2 with sensible defaults."""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
return ConversationV2(
|
||||
conversation_id=conversation_id,
|
||||
platform_id=platform_id,
|
||||
user_id=user_id,
|
||||
title=title,
|
||||
persona_id=persona_id,
|
||||
content=content or [],
|
||||
token_usage=token_usage,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_db():
|
||||
"""Create a mock BaseDatabase."""
|
||||
db = MagicMock()
|
||||
db.create_conversation = AsyncMock()
|
||||
db.get_conversation_by_id = AsyncMock()
|
||||
db.get_conversations = AsyncMock()
|
||||
db.get_filtered_conversations = AsyncMock()
|
||||
db.delete_conversation = AsyncMock()
|
||||
db.delete_conversations_by_user_id = AsyncMock()
|
||||
db.update_conversation = AsyncMock()
|
||||
return db
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mgr(mock_db):
|
||||
"""Create a ConversationManager with mocked database."""
|
||||
return ConversationManager(mock_db)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConversationManagerInit:
|
||||
"""Tests for ConversationManager.__init__()."""
|
||||
|
||||
def test_init_sets_attributes(self, mock_db):
|
||||
"""Verify __init__ sets all initial attributes."""
|
||||
mgr = ConversationManager(mock_db)
|
||||
assert mgr.db is mock_db
|
||||
assert mgr.session_conversations == {}
|
||||
assert mgr.save_interval == 60
|
||||
assert mgr._on_session_deleted_callbacks == []
|
||||
|
||||
def test_register_on_session_deleted(self, mgr):
|
||||
"""Verify register_on_session_deleted adds callback."""
|
||||
async def callback(_umo: str):
|
||||
pass
|
||||
|
||||
mgr.register_on_session_deleted(callback)
|
||||
assert callback in mgr._on_session_deleted_callbacks
|
||||
|
||||
|
||||
class TestConversationManagerConvertV2ToV1:
|
||||
"""Tests for _convert_conv_from_v2_to_v1()."""
|
||||
|
||||
def test_converts_v2_to_v1(self, mgr):
|
||||
"""Verify basic conversion of ConversationV2 to Conversation."""
|
||||
conv_v2 = _make_conv_v2(
|
||||
conversation_id="c1",
|
||||
platform_id="qq",
|
||||
user_id="user:123",
|
||||
title="Chat",
|
||||
persona_id="p1",
|
||||
content=[{"role": "user", "content": "hi"}],
|
||||
token_usage=100,
|
||||
)
|
||||
conv = mgr._convert_conv_from_v2_to_v1(conv_v2)
|
||||
assert isinstance(conv, Conversation)
|
||||
assert conv.cid == "c1"
|
||||
assert conv.platform_id == "qq"
|
||||
assert conv.user_id == "user:123"
|
||||
assert conv.title == "Chat"
|
||||
assert conv.persona_id == "p1"
|
||||
assert conv.token_usage == 100
|
||||
assert json.loads(conv.history) == [{"role": "user", "content": "hi"}]
|
||||
|
||||
def test_converts_empty_content(self, mgr):
|
||||
"""Verify conversion handles None content."""
|
||||
conv_v2 = _make_conv_v2(content=None)
|
||||
conv = mgr._convert_conv_from_v2_to_v1(conv_v2)
|
||||
assert json.loads(conv.history) == []
|
||||
|
||||
def test_converts_null_timestamps(self, mgr):
|
||||
"""Verify conversion handles None timestamps."""
|
||||
conv_v2 = _make_conv_v2()
|
||||
conv_v2.created_at = None
|
||||
conv_v2.updated_at = None
|
||||
conv = mgr._convert_conv_from_v2_to_v1(conv_v2)
|
||||
assert conv.created_at == 0
|
||||
assert conv.updated_at == 0
|
||||
|
||||
|
||||
class TestConversationManagerNewConversation:
|
||||
"""Tests for new_conversation()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_conversation_creates_and_caches(self, mgr, mock_db):
|
||||
"""Verify new_conversation creates via DB and caches the ID."""
|
||||
created_conv = _make_conv_v2(conversation_id="new-conv")
|
||||
mock_db.create_conversation.return_value = created_conv
|
||||
|
||||
with patch("astrbot.core.conversation_mgr.sp.session_put", AsyncMock()) as sp_put:
|
||||
cid = await mgr.new_conversation(
|
||||
"qq:group:123",
|
||||
platform_id="qq",
|
||||
content=[{"role": "user", "content": "hi"}],
|
||||
title="New Chat",
|
||||
persona_id="p1",
|
||||
)
|
||||
|
||||
assert cid == "new-conv"
|
||||
assert mgr.session_conversations["qq:group:123"] == "new-conv"
|
||||
mock_db.create_conversation.assert_awaited_once_with(
|
||||
user_id="qq:group:123",
|
||||
platform_id="qq",
|
||||
content=[{"role": "user", "content": "hi"}],
|
||||
title="New Chat",
|
||||
persona_id="p1",
|
||||
)
|
||||
sp_put.assert_awaited_once_with("qq:group:123", "sel_conv_id", "new-conv")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_conversation_infers_platform_id(self, mgr, mock_db):
|
||||
"""Verify platform_id is inferred from unified_msg_origin when not provided."""
|
||||
created_conv = _make_conv_v2(conversation_id="cid")
|
||||
mock_db.create_conversation.return_value = created_conv
|
||||
|
||||
with patch("astrbot.core.conversation_mgr.sp.session_put", AsyncMock()):
|
||||
cid = await mgr.new_conversation("discord:dm:456")
|
||||
|
||||
assert cid == "cid"
|
||||
mock_db.create_conversation.assert_awaited_once_with(
|
||||
user_id="discord:dm:456",
|
||||
platform_id="discord",
|
||||
content=None,
|
||||
title=None,
|
||||
persona_id=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_conversation_fallback_platform(self, mgr, mock_db):
|
||||
"""Verify platform_id falls back to 'unknown' when it cannot be inferred."""
|
||||
created_conv = _make_conv_v2(conversation_id="cid")
|
||||
mock_db.create_conversation.return_value = created_conv
|
||||
|
||||
with patch("astrbot.core.conversation_mgr.sp.session_put", AsyncMock()):
|
||||
cid = await mgr.new_conversation("short")
|
||||
|
||||
assert cid == "cid"
|
||||
mock_db.create_conversation.assert_awaited_once_with(
|
||||
user_id="short",
|
||||
platform_id="unknown",
|
||||
content=None,
|
||||
title=None,
|
||||
persona_id=None,
|
||||
)
|
||||
|
||||
|
||||
class TestConversationManagerSwitchConversation:
|
||||
"""Tests for switch_conversation()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_switch_conversation(self, mgr):
|
||||
"""Verify switch updates cache and persists to session prefs."""
|
||||
mgr.session_conversations["origin:1"] = "old-id"
|
||||
|
||||
with patch("astrbot.core.conversation_mgr.sp.session_put", AsyncMock()) as sp_put:
|
||||
await mgr.switch_conversation("origin:1", "new-id")
|
||||
|
||||
assert mgr.session_conversations["origin:1"] == "new-id"
|
||||
sp_put.assert_awaited_once_with("origin:1", "sel_conv_id", "new-id")
|
||||
|
||||
|
||||
class TestConversationManagerDeleteConversation:
|
||||
"""Tests for delete_conversation()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_current_conversation(self, mgr, mock_db):
|
||||
"""Verify deleting the current conversation clears cache."""
|
||||
mgr.session_conversations["origin:1"] = "conv-id"
|
||||
mock_db.get_conversation_by_id = AsyncMock(return_value=None)
|
||||
mock_db.delete_conversation = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.conversation_mgr.sp.session_remove",
|
||||
AsyncMock(),
|
||||
) as sp_remove,
|
||||
patch.object(
|
||||
mgr,
|
||||
"get_curr_conversation_id",
|
||||
AsyncMock(return_value="conv-id"),
|
||||
),
|
||||
):
|
||||
await mgr.delete_conversation("origin:1")
|
||||
|
||||
mock_db.delete_conversation.assert_awaited_once_with(cid="conv-id")
|
||||
assert "origin:1" not in mgr.session_conversations
|
||||
sp_remove.assert_awaited_once_with("origin:1", "sel_conv_id")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_specific_conversation(self, mgr, mock_db):
|
||||
"""Verify deleting a non-current conversation by ID."""
|
||||
mgr.session_conversations["origin:1"] = "current-id"
|
||||
mock_db.delete_conversation = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.conversation_mgr.sp.session_remove",
|
||||
AsyncMock(),
|
||||
) as sp_remove,
|
||||
patch.object(
|
||||
mgr,
|
||||
"get_curr_conversation_id",
|
||||
AsyncMock(return_value="current-id"),
|
||||
),
|
||||
):
|
||||
await mgr.delete_conversation("origin:1", conversation_id="other-id")
|
||||
|
||||
mock_db.delete_conversation.assert_awaited_once_with(cid="other-id")
|
||||
# Current ID should NOT be removed since it differs
|
||||
assert mgr.session_conversations.get("origin:1") == "current-id"
|
||||
sp_remove.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_no_conversation_id_fallback(self, mgr, mock_db):
|
||||
"""Verify when no conv ID is given and none cached, do nothing."""
|
||||
mgr.session_conversations = {}
|
||||
mock_db.delete_conversation = AsyncMock()
|
||||
|
||||
with patch(
|
||||
"astrbot.core.conversation_mgr.sp.session_remove",
|
||||
AsyncMock(),
|
||||
):
|
||||
await mgr.delete_conversation("origin:1")
|
||||
|
||||
mock_db.delete_conversation.assert_not_called()
|
||||
|
||||
|
||||
class TestConversationManagerDeleteConversationsByUserId:
|
||||
"""Tests for delete_conversations_by_user_id()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_all_and_trigger_callbacks(self, mgr, mock_db):
|
||||
"""Verify deleting all conversations cleans cache and triggers callbacks."""
|
||||
mgr.session_conversations["origin:1"] = "c1"
|
||||
mock_db.delete_conversations_by_user_id = AsyncMock()
|
||||
|
||||
callback = AsyncMock()
|
||||
mgr.register_on_session_deleted(callback)
|
||||
|
||||
with patch(
|
||||
"astrbot.core.conversation_mgr.sp.session_remove",
|
||||
AsyncMock(),
|
||||
) as sp_remove:
|
||||
await mgr.delete_conversations_by_user_id("origin:1")
|
||||
|
||||
mock_db.delete_conversations_by_user_id.assert_awaited_once_with(
|
||||
user_id="origin:1",
|
||||
)
|
||||
assert "origin:1" not in mgr.session_conversations
|
||||
sp_remove.assert_awaited_once_with("origin:1", "sel_conv_id")
|
||||
callback.assert_awaited_once_with("origin:1")
|
||||
|
||||
|
||||
class TestConversationManagerGetCurrConversationId:
|
||||
"""Tests for get_curr_conversation_id()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_cached_value(self, mgr):
|
||||
"""Verify returns cached conversation ID without hitting session prefs."""
|
||||
mgr.session_conversations["origin:1"] = "cached-id"
|
||||
|
||||
with patch("astrbot.core.conversation_mgr.sp.session_get", AsyncMock()) as sp_get:
|
||||
cid = await mgr.get_curr_conversation_id("origin:1")
|
||||
|
||||
assert cid == "cached-id"
|
||||
sp_get.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetches_from_session_prefs_and_caches(self, mgr):
|
||||
"""Verify fetches from session prefs when not cached, then caches it."""
|
||||
mgr.session_conversations = {}
|
||||
|
||||
with patch(
|
||||
"astrbot.core.conversation_mgr.sp.session_get",
|
||||
AsyncMock(return_value="pref-id"),
|
||||
) as sp_get:
|
||||
cid = await mgr.get_curr_conversation_id("origin:1")
|
||||
|
||||
assert cid == "pref-id"
|
||||
assert mgr.session_conversations["origin:1"] == "pref-id"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_none_when_not_found(self, mgr):
|
||||
"""Verify returns None when no conversation is known."""
|
||||
mgr.session_conversations = {}
|
||||
|
||||
with patch(
|
||||
"astrbot.core.conversation_mgr.sp.session_get",
|
||||
AsyncMock(return_value=None),
|
||||
):
|
||||
cid = await mgr.get_curr_conversation_id("origin:1")
|
||||
|
||||
assert cid is None
|
||||
|
||||
|
||||
class TestConversationManagerGetConversation:
|
||||
"""Tests for get_conversation()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_existing_conversation(self, mgr, mock_db):
|
||||
"""Verify retrieves and converts existing conversation."""
|
||||
conv_v2 = _make_conv_v2(conversation_id="c1")
|
||||
mock_db.get_conversation_by_id.return_value = conv_v2
|
||||
|
||||
conv = await mgr.get_conversation("origin:1", "c1")
|
||||
|
||||
assert conv is not None
|
||||
assert conv.cid == "c1"
|
||||
mock_db.get_conversation_by_id.assert_awaited_once_with(cid="c1")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_non_existing_not_created(self, mgr, mock_db):
|
||||
"""Verify returns None when not found and create_if_not_exists is False."""
|
||||
mock_db.get_conversation_by_id.return_value = None
|
||||
|
||||
conv = await mgr.get_conversation("origin:1", "nonexistent")
|
||||
|
||||
assert conv is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_non_existing_creates_new(self, mgr, mock_db):
|
||||
"""Verify creates new conversation when not found and create_if_not_exists is True."""
|
||||
mock_db.get_conversation_by_id.side_effect = [
|
||||
None,
|
||||
_make_conv_v2(conversation_id="new-c1"),
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
mgr,
|
||||
"new_conversation",
|
||||
AsyncMock(return_value="new-c1"),
|
||||
) as mock_new:
|
||||
conv = await mgr.get_conversation(
|
||||
"origin:1",
|
||||
"nonexistent",
|
||||
create_if_not_exists=True,
|
||||
)
|
||||
|
||||
assert conv is not None
|
||||
assert conv.cid == "new-c1"
|
||||
mock_new.assert_awaited_once_with("origin:1")
|
||||
|
||||
|
||||
class TestConversationManagerGetConversations:
|
||||
"""Tests for get_conversations()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_conversations(self, mgr, mock_db):
|
||||
"""Verify retrieving multiple conversations."""
|
||||
convs_v2 = [
|
||||
_make_conv_v2(conversation_id="c1"),
|
||||
_make_conv_v2(conversation_id="c2"),
|
||||
]
|
||||
mock_db.get_conversations.return_value = convs_v2
|
||||
|
||||
convs = await mgr.get_conversations(unified_msg_origin="origin:1")
|
||||
|
||||
assert len(convs) == 2
|
||||
assert convs[0].cid == "c1"
|
||||
assert convs[1].cid == "c2"
|
||||
mock_db.get_conversations.assert_awaited_once_with(
|
||||
user_id="origin:1",
|
||||
platform_id=None,
|
||||
)
|
||||
|
||||
|
||||
class TestConversationManagerGetFilteredConversations:
|
||||
"""Tests for get_filtered_conversations()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_filtered_conversations(self, mgr, mock_db):
|
||||
"""Verify filtered conversation retrieval."""
|
||||
convs_v2 = [_make_conv_v2(conversation_id="c1")]
|
||||
mock_db.get_filtered_conversations.return_value = (convs_v2, 1)
|
||||
|
||||
convs, cnt = await mgr.get_filtered_conversations(
|
||||
page=1,
|
||||
page_size=20,
|
||||
platform_ids=["qq"],
|
||||
search_query="test",
|
||||
)
|
||||
|
||||
assert len(convs) == 1
|
||||
assert convs[0].cid == "c1"
|
||||
assert cnt == 1
|
||||
mock_db.get_filtered_conversations.assert_awaited_once_with(
|
||||
page=1,
|
||||
page_size=20,
|
||||
platform_ids=["qq"],
|
||||
search_query="test",
|
||||
)
|
||||
|
||||
|
||||
class TestConversationManagerUpdateConversation:
|
||||
"""Tests for update_conversation()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_without_id_uses_current(self, mgr, mock_db):
|
||||
"""Verify update uses current conversation ID when not provided."""
|
||||
with patch.object(
|
||||
mgr,
|
||||
"get_curr_conversation_id",
|
||||
AsyncMock(return_value="current-id"),
|
||||
):
|
||||
await mgr.update_conversation(
|
||||
"origin:1",
|
||||
history=[{"role": "user", "content": "hi"}],
|
||||
title="Updated",
|
||||
persona_id="p2",
|
||||
token_usage=50,
|
||||
)
|
||||
|
||||
mock_db.update_conversation.assert_awaited_once_with(
|
||||
cid="current-id",
|
||||
title="Updated",
|
||||
persona_id="p2",
|
||||
clear_persona=False,
|
||||
content=[{"role": "user", "content": "hi"}],
|
||||
token_usage=50,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_with_id(self, mgr, mock_db):
|
||||
"""Verify update with explicit conversation ID."""
|
||||
await mgr.update_conversation(
|
||||
"origin:1",
|
||||
conversation_id="explicit-id",
|
||||
title="New Title",
|
||||
)
|
||||
|
||||
mock_db.update_conversation.assert_awaited_once_with(
|
||||
cid="explicit-id",
|
||||
title="New Title",
|
||||
persona_id=None,
|
||||
clear_persona=False,
|
||||
content=None,
|
||||
token_usage=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_without_id_no_current_does_nothing(self, mgr, mock_db):
|
||||
"""Verify update does nothing when no ID is available."""
|
||||
with patch.object(
|
||||
mgr,
|
||||
"get_curr_conversation_id",
|
||||
AsyncMock(return_value=None),
|
||||
):
|
||||
await mgr.update_conversation("origin:1", title="New Title")
|
||||
|
||||
mock_db.update_conversation.assert_not_called()
|
||||
|
||||
|
||||
class TestConversationManagerUpdateConversationTitle:
|
||||
"""Tests for update_conversation_title()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_title_delegates(self, mgr):
|
||||
"""Verify update_conversation_title delegates to update_conversation."""
|
||||
with patch.object(mgr, "update_conversation", AsyncMock()) as mock_update:
|
||||
await mgr.update_conversation_title("origin:1", "New Title", "conv-id")
|
||||
|
||||
mock_update.assert_awaited_once_with(
|
||||
unified_msg_origin="origin:1",
|
||||
conversation_id="conv-id",
|
||||
title="New Title",
|
||||
)
|
||||
|
||||
|
||||
class TestConversationManagerUpdateConversationPersonaId:
|
||||
"""Tests for update_conversation_persona_id()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_persona_delegates(self, mgr):
|
||||
"""Verify update_conversation_persona_id delegates to update_conversation."""
|
||||
with patch.object(mgr, "update_conversation", AsyncMock()) as mock_update:
|
||||
await mgr.update_conversation_persona_id(
|
||||
"origin:1",
|
||||
"persona-123",
|
||||
"conv-id",
|
||||
)
|
||||
|
||||
mock_update.assert_awaited_once_with(
|
||||
unified_msg_origin="origin:1",
|
||||
conversation_id="conv-id",
|
||||
persona_id="persona-123",
|
||||
)
|
||||
|
||||
|
||||
class TestConversationManagerUnsetConversationPersona:
|
||||
"""Tests for unset_conversation_persona()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unset_persona_delegates_with_clear(self, mgr):
|
||||
"""Verify unset_conversation_persona delegates with clear_persona=True."""
|
||||
with patch.object(mgr, "update_conversation", AsyncMock()) as mock_update:
|
||||
await mgr.unset_conversation_persona("origin:1", "conv-id")
|
||||
|
||||
mock_update.assert_awaited_once_with(
|
||||
unified_msg_origin="origin:1",
|
||||
conversation_id="conv-id",
|
||||
clear_persona=True,
|
||||
)
|
||||
|
||||
|
||||
class TestConversationManagerAddMessagePair:
|
||||
"""Tests for add_message_pair()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_message_pair_dicts(self, mgr, mock_db):
|
||||
"""Verify adding message pair using plain dicts."""
|
||||
conv_v2 = _make_conv_v2(content=[])
|
||||
mock_db.get_conversation_by_id.return_value = conv_v2
|
||||
|
||||
await mgr.add_message_pair(
|
||||
"c1",
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "hi there"},
|
||||
)
|
||||
|
||||
mock_db.update_conversation.assert_awaited_once_with(
|
||||
cid="c1",
|
||||
content=[
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "hi there"},
|
||||
],
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_message_pair_appends_to_existing(self, mgr, mock_db):
|
||||
"""Verify adding message pair appends to existing history."""
|
||||
conv_v2 = _make_conv_v2(content=[{"role": "user", "content": "prev"}])
|
||||
mock_db.get_conversation_by_id.return_value = conv_v2
|
||||
|
||||
await mgr.add_message_pair(
|
||||
"c1",
|
||||
{"role": "user", "content": "q2"},
|
||||
{"role": "assistant", "content": "a2"},
|
||||
)
|
||||
|
||||
mock_db.update_conversation.assert_awaited_once_with(
|
||||
cid="c1",
|
||||
content=[
|
||||
{"role": "user", "content": "prev"},
|
||||
{"role": "user", "content": "q2"},
|
||||
{"role": "assistant", "content": "a2"},
|
||||
],
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_message_pair_conv_not_found(self, mgr, mock_db):
|
||||
"""Verify raises when conversation is not found."""
|
||||
mock_db.get_conversation_by_id.return_value = None
|
||||
|
||||
with pytest.raises(Exception, match="Conversation with id nonexistent not found"):
|
||||
await mgr.add_message_pair(
|
||||
"nonexistent",
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "hello"},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_message_pair_segment_objects(self, mgr, mock_db):
|
||||
"""Verify adding message pair using UserMessageSegment / AssistantMessageSegment."""
|
||||
from astrbot.core.agent.message import (
|
||||
AssistantMessageSegment,
|
||||
UserMessageSegment,
|
||||
)
|
||||
|
||||
conv_v2 = _make_conv_v2(content=[])
|
||||
mock_db.get_conversation_by_id.return_value = conv_v2
|
||||
|
||||
user_msg = UserMessageSegment(content="hello")
|
||||
assistant_msg = AssistantMessageSegment(content="world")
|
||||
|
||||
with (
|
||||
patch.object(user_msg, "model_dump", return_value={"role": "user", "content": "hello"}),
|
||||
patch.object(
|
||||
assistant_msg,
|
||||
"model_dump",
|
||||
return_value={"role": "assistant", "content": "world"},
|
||||
),
|
||||
):
|
||||
await mgr.add_message_pair("c1", user_msg, assistant_msg)
|
||||
|
||||
mock_db.update_conversation.assert_awaited_once_with(
|
||||
cid="c1",
|
||||
content=[
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "world"},
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
class TestConversationManagerGetHumanReadableContext:
|
||||
"""Tests for get_human_readable_context()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_context(self, mgr):
|
||||
"""Verify basic context formatting."""
|
||||
mgr._convert_conv_from_v2_to_v1 = MagicMock(
|
||||
return_value=Conversation(
|
||||
platform_id="qq",
|
||||
user_id="user:1",
|
||||
cid="c1",
|
||||
history=json.dumps([
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "world"},
|
||||
{"role": "user", "content": "how are you"},
|
||||
{"role": "assistant", "content": "fine"},
|
||||
]),
|
||||
title="Chat",
|
||||
created_at=1000,
|
||||
updated_at=1001,
|
||||
),
|
||||
)
|
||||
|
||||
contexts, total_pages = await mgr.get_human_readable_context(
|
||||
"origin:1",
|
||||
"c1",
|
||||
page=1,
|
||||
page_size=10,
|
||||
)
|
||||
|
||||
# Order: most recent pair first
|
||||
assert contexts == [
|
||||
"User: how are you",
|
||||
"Assistant: fine",
|
||||
"User: hello",
|
||||
"Assistant: world",
|
||||
]
|
||||
assert total_pages == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_context_no_conversation(self, mgr):
|
||||
"""Verify returns empty when conversation not found."""
|
||||
with patch.object(
|
||||
mgr,
|
||||
"get_conversation",
|
||||
AsyncMock(return_value=None),
|
||||
):
|
||||
contexts, total_pages = await mgr.get_human_readable_context(
|
||||
"origin:1",
|
||||
"nonexistent",
|
||||
)
|
||||
|
||||
assert contexts == []
|
||||
assert total_pages == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_context_pagination(self, mgr):
|
||||
"""Verify context pagination."""
|
||||
records = []
|
||||
for i in range(5):
|
||||
records.append({"role": "user", "content": f"q{i}"})
|
||||
records.append({"role": "assistant", "content": f"a{i}"})
|
||||
|
||||
mgr._convert_conv_from_v2_to_v1 = MagicMock(
|
||||
return_value=Conversation(
|
||||
platform_id="qq",
|
||||
user_id="user:1",
|
||||
cid="c1",
|
||||
history=json.dumps(records),
|
||||
title="Chat",
|
||||
created_at=1000,
|
||||
updated_at=1001,
|
||||
),
|
||||
)
|
||||
|
||||
contexts, total_pages = await mgr.get_human_readable_context(
|
||||
"origin:1",
|
||||
"c1",
|
||||
page=1,
|
||||
page_size=2,
|
||||
)
|
||||
|
||||
# 5 pairs = 10 records, reversed, page_size=2
|
||||
assert len(contexts) == 2
|
||||
assert total_pages == 5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_context_with_tool_calls(self, mgr):
|
||||
"""Verify context handles tool_calls in assistant messages."""
|
||||
mgr._convert_conv_from_v2_to_v1 = MagicMock(
|
||||
return_value=Conversation(
|
||||
platform_id="qq",
|
||||
user_id="user:1",
|
||||
cid="c1",
|
||||
history=json.dumps([
|
||||
{"role": "user", "content": "search weather"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [{"function": {"name": "get_weather"}}],
|
||||
},
|
||||
]),
|
||||
title="Chat",
|
||||
created_at=1000,
|
||||
updated_at=1001,
|
||||
),
|
||||
)
|
||||
|
||||
contexts, _ = await mgr.get_human_readable_context("origin:1", "c1")
|
||||
|
||||
assert "Assistant: [函数调用]" in contexts[0]
|
||||
|
||||
|
||||
class TestConversationManagerTriggerSessionDeleted:
|
||||
"""Tests for _trigger_session_deleted()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_triggers_callbacks(self, mgr):
|
||||
"""Verify all callbacks are triggered."""
|
||||
cb1 = AsyncMock()
|
||||
cb2 = AsyncMock()
|
||||
mgr.register_on_session_deleted(cb1)
|
||||
mgr.register_on_session_deleted(cb2)
|
||||
|
||||
await mgr._trigger_session_deleted("origin:1")
|
||||
|
||||
cb1.assert_awaited_once_with("origin:1")
|
||||
cb2.assert_awaited_once_with("origin:1")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_error_does_not_block_others(self, mgr):
|
||||
"""Verify one failing callback does not prevent others from running."""
|
||||
cb1 = AsyncMock(side_effect=RuntimeError("fail"))
|
||||
cb2 = AsyncMock()
|
||||
mgr.register_on_session_deleted(cb1)
|
||||
mgr.register_on_session_deleted(cb2)
|
||||
|
||||
# Should not raise
|
||||
await mgr._trigger_session_deleted("origin:1")
|
||||
|
||||
cb2.assert_awaited_once_with("origin:1")
|
||||
@@ -0,0 +1,242 @@
|
||||
"""Tests for CronMessageEvent."""
|
||||
|
||||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.cron.events import CronMessageEvent
|
||||
from astrbot.core.message.components import Plain
|
||||
from astrbot.core.message.message_event_result import MessageChain
|
||||
from astrbot.core.platform.astr_message_event import AstrMessageEvent
|
||||
from astrbot.core.platform.message_session import MessageSession
|
||||
from astrbot.core.platform.message_type import MessageType
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_context():
|
||||
ctx = MagicMock()
|
||||
ctx.send_message = AsyncMock()
|
||||
return ctx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session():
|
||||
session = MagicMock(spec=MessageSession)
|
||||
session.session_id = "test-session-id"
|
||||
session.platform_id = "test-platform"
|
||||
return session
|
||||
|
||||
|
||||
class TestCronMessageEventInit:
|
||||
"""Tests for CronMessageEvent construction."""
|
||||
|
||||
def test_init_default_params(self, mock_context, mock_session):
|
||||
"""Default params produce a synthetic event with correct attributes."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Hello from cron",
|
||||
)
|
||||
assert event.message_str == "Hello from cron"
|
||||
assert event.is_at_or_wake_command is True
|
||||
assert event.is_wake is True
|
||||
assert event.session == mock_session
|
||||
assert event.context_obj == mock_context
|
||||
|
||||
def test_init_sender_defaults(self, mock_context, mock_session):
|
||||
"""Default sender_id and sender_name are 'astrbot' and 'Scheduler'."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
assert event.message_obj.self_id == "astrbot"
|
||||
assert event.message_obj.sender.nickname == "Scheduler"
|
||||
assert event.message_obj.sender.user_id == "test-session-id"
|
||||
|
||||
def test_init_custom_sender(self, mock_context, mock_session):
|
||||
"""Custom sender_id and sender_name are reflected on the message object."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
sender_id="custom-bot",
|
||||
sender_name="CustomName",
|
||||
)
|
||||
assert event.message_obj.self_id == "custom-bot"
|
||||
assert event.message_obj.sender.nickname == "CustomName"
|
||||
|
||||
def test_init_with_extras(self, mock_context, mock_session):
|
||||
"""Extras dict is merged into _extras."""
|
||||
extras = {"origin": "api", "cron_payload": {"key": "value"}}
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
extras=extras,
|
||||
)
|
||||
assert event._extras.get("origin") == "api"
|
||||
assert event._extras["cron_payload"] == {"key": "value"}
|
||||
|
||||
def test_init_group_message_type(self, mock_context, mock_session):
|
||||
"""MessageType.GROUP_MESSAGE is preserved in the message object."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
message_type=MessageType.GROUP_MESSAGE,
|
||||
)
|
||||
assert event.message_obj.type == MessageType.GROUP_MESSAGE
|
||||
|
||||
def test_init_platform_meta(self, mock_context, mock_session):
|
||||
"""Platform metadata is set to cron / CronJob / platform_id."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
assert event.platform_meta.name == "cron"
|
||||
assert event.platform_meta.description == "CronJob"
|
||||
assert event.platform_meta.id == "test-platform"
|
||||
|
||||
def test_init_message_components(self, mock_context, mock_session):
|
||||
"""The message is wrapped in a Plain component."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Cron message text",
|
||||
)
|
||||
assert len(event.message_obj.message) == 1
|
||||
assert isinstance(event.message_obj.message[0], Plain)
|
||||
assert event.message_obj.message[0].text == "Cron message text"
|
||||
assert event.message_obj.message_str == "Cron message text"
|
||||
assert event.message_obj.raw_message == "Cron message text"
|
||||
|
||||
def test_init_timestamp_within_range(self, mock_context, mock_session):
|
||||
"""Timestamp is set to the current time."""
|
||||
before = int(time.time())
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
after = int(time.time())
|
||||
assert before <= event.message_obj.timestamp <= after
|
||||
|
||||
def test_init_message_id_is_uuid_hex(self, mock_context, mock_session):
|
||||
"""Message ID is a 32-char hex string (uuid4 hex)."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
mid = event.message_obj.message_id
|
||||
assert isinstance(mid, str)
|
||||
assert len(mid) == 32
|
||||
# All characters should be valid hex digits
|
||||
int(mid, 16)
|
||||
|
||||
def test_init_none_extras_does_not_raise(self, mock_context, mock_session):
|
||||
"""Passing extras=None does not raise."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
extras=None,
|
||||
)
|
||||
# Should not raise; _extras remains default
|
||||
assert isinstance(event._extras, dict)
|
||||
|
||||
def test_init_uses_session_session_id(self, mock_context, mock_session):
|
||||
"""Session ID is taken from session.session_id."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
assert event.session_id == "test-session-id"
|
||||
|
||||
|
||||
class TestCronMessageEventSend:
|
||||
"""Tests for CronMessageEvent.send."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_calls_context(self, mock_context, mock_session):
|
||||
"""send delegates to context.send_message on the original session."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
chain = MessageChain([Plain("response")])
|
||||
|
||||
with patch.object(AstrMessageEvent, "send", new_callable=AsyncMock):
|
||||
await event.send(chain)
|
||||
|
||||
mock_context.send_message.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_none_is_noop(self, mock_context, mock_session):
|
||||
"""send(None) returns immediately without calling context."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
with patch.object(AstrMessageEvent, "send", new_callable=AsyncMock) as mock_super:
|
||||
await event.send(None)
|
||||
|
||||
mock_context.send_message.assert_not_called()
|
||||
mock_super.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_calls_super(self, mock_context, mock_session):
|
||||
"""send also delegates to super().send()."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
chain = MessageChain([Plain("response")])
|
||||
|
||||
with patch.object(AstrMessageEvent, "send", new_callable=AsyncMock) as mock_super:
|
||||
await event.send(chain)
|
||||
|
||||
mock_super.assert_awaited_once_with(chain)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_streaming_iterates_generator(self, mock_context, mock_session):
|
||||
"""send_streaming calls send for each yielded chain."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
|
||||
async def gen():
|
||||
yield MessageChain([Plain("part1")])
|
||||
yield MessageChain([Plain("part2")])
|
||||
|
||||
with patch.object(event, "send", new_callable=AsyncMock) as mock_send:
|
||||
await event.send_streaming(gen())
|
||||
|
||||
assert mock_send.call_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_streaming_empty_generator(self, mock_context, mock_session):
|
||||
"""send_streaming does not call send when the generator is empty."""
|
||||
event = CronMessageEvent(
|
||||
context=mock_context,
|
||||
session=mock_session,
|
||||
message="Test",
|
||||
)
|
||||
|
||||
async def empty_gen():
|
||||
if False:
|
||||
yield # pragma: no cover
|
||||
|
||||
with patch.object(event, "send", new_callable=AsyncMock) as mock_send:
|
||||
await event.send_streaming(empty_gen())
|
||||
|
||||
mock_send.assert_not_called()
|
||||
@@ -0,0 +1,428 @@
|
||||
"""Supplementary edge-case tests for CronJobManager.
|
||||
|
||||
Covers validation branches, interval-based jobs, run-once auto-cleanup,
|
||||
and failure paths not covered by the main test suite.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.cron.manager import CronJobManager, CronJobSchedulingError
|
||||
from astrbot.core.db.po import CronJob
|
||||
|
||||
|
||||
# ---- Fixtures (self-contained) ----
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_db():
|
||||
db = MagicMock()
|
||||
db.create_cron_job = AsyncMock()
|
||||
db.get_cron_job = AsyncMock()
|
||||
db.update_cron_job = AsyncMock()
|
||||
db.delete_cron_job = AsyncMock()
|
||||
db.list_cron_jobs = AsyncMock(return_value=[])
|
||||
return db
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_context():
|
||||
ctx = MagicMock()
|
||||
ctx.get_config = MagicMock(return_value={"admins_id": []})
|
||||
ctx.conversation_manager = MagicMock()
|
||||
return ctx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def cron_manager(mock_db):
|
||||
return CronJobManager(mock_db)
|
||||
|
||||
|
||||
# ---- add_basic_job validation edge cases ----
|
||||
|
||||
|
||||
class TestAddBasicJobEdgeCases:
|
||||
"""Validation and interval-based variants for add_basic_job."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_basic_job_with_interval_seconds(self, cron_manager, mock_db):
|
||||
"""interval_seconds can be used instead of cron_expression."""
|
||||
cron_manager._started = True
|
||||
job = CronJob(job_id="interval-job", name="Interval", job_type="basic")
|
||||
mock_db.create_cron_job.return_value = job
|
||||
|
||||
handler = MagicMock()
|
||||
result = await cron_manager.add_basic_job(
|
||||
name="Interval",
|
||||
interval_seconds=300,
|
||||
handler=handler,
|
||||
description="Every 5 min",
|
||||
)
|
||||
|
||||
assert result == job
|
||||
mock_db.create_cron_job.assert_called_once()
|
||||
# interval should be packed into the payload
|
||||
call_payload = mock_db.create_cron_job.call_args.kwargs["payload"]
|
||||
assert "interval_seconds" in call_payload
|
||||
assert call_payload["interval_seconds"] == 300
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_basic_job_both_cron_and_interval_raises(self, cron_manager, mock_db):
|
||||
"""Providing both cron_expression and interval_seconds raises ValueError."""
|
||||
handler = MagicMock()
|
||||
with pytest.raises(ValueError, match="must have exactly one value"):
|
||||
await cron_manager.add_basic_job(
|
||||
name="Bad",
|
||||
cron_expression="0 9 * * *",
|
||||
interval_seconds=300,
|
||||
handler=handler,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_basic_job_neither_cron_nor_interval_raises(self, cron_manager, mock_db):
|
||||
"""Providing neither cron_expression nor interval_seconds raises ValueError."""
|
||||
handler = MagicMock()
|
||||
with pytest.raises(ValueError, match="must have exactly one value"):
|
||||
await cron_manager.add_basic_job(
|
||||
name="Bad",
|
||||
handler=handler,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_basic_job_payload_passed_through(self, cron_manager, mock_db):
|
||||
"""User-supplied payload is forwarded to create_cron_job."""
|
||||
cron_manager._started = True
|
||||
job = CronJob(job_id="payload-job", name="Payload", job_type="basic")
|
||||
mock_db.create_cron_job.return_value = job
|
||||
|
||||
handler = MagicMock()
|
||||
user_payload = {"custom_key": "custom_value"}
|
||||
await cron_manager.add_basic_job(
|
||||
name="Payload",
|
||||
cron_expression="0 9 * * *",
|
||||
handler=handler,
|
||||
payload=user_payload,
|
||||
)
|
||||
|
||||
call_payload = mock_db.create_cron_job.call_args.kwargs["payload"]
|
||||
assert call_payload["custom_key"] == "custom_value"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_basic_job_non_persistent_not_scheduled(self, cron_manager, mock_db):
|
||||
"""A disabled non-persistent job is stored but not scheduled."""
|
||||
job = CronJob(
|
||||
job_id="np-job", name="NonPersist", job_type="basic", enabled=False
|
||||
)
|
||||
mock_db.create_cron_job.return_value = job
|
||||
|
||||
handler = MagicMock()
|
||||
with patch.object(cron_manager, "_schedule_job") as mock_schedule:
|
||||
result = await cron_manager.add_basic_job(
|
||||
name="NonPersist",
|
||||
cron_expression="0 9 * * *",
|
||||
handler=handler,
|
||||
persistent=False,
|
||||
enabled=False,
|
||||
)
|
||||
|
||||
assert result == job
|
||||
mock_schedule.assert_not_called()
|
||||
|
||||
|
||||
# ---- _run_job full lifecycle edge cases ----
|
||||
|
||||
|
||||
class TestRunJobEdgeCases:
|
||||
"""Full lifecycle for _run_job covering status transitions and cleanup."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_job_basic_completed(self, cron_manager, mock_db):
|
||||
"""A basic job transitions to 'completed' after successful run."""
|
||||
job = CronJob(
|
||||
job_id="basic-done",
|
||||
name="Done",
|
||||
job_type="basic",
|
||||
enabled=True,
|
||||
cron_expression="0 9 * * *",
|
||||
)
|
||||
mock_db.get_cron_job.return_value = job
|
||||
|
||||
handler = MagicMock(return_value=None)
|
||||
cron_manager._basic_handlers["basic-done"] = handler
|
||||
|
||||
await cron_manager._run_job("basic-done")
|
||||
|
||||
# Status should be updated twice: running -> completed
|
||||
update_calls = [
|
||||
c for c in mock_db.update_cron_job.call_args_list
|
||||
]
|
||||
# At minimum one call with status="running" and one with status="completed"
|
||||
running_calls = [
|
||||
c for c in update_calls
|
||||
if c.kwargs.get("status") == "running"
|
||||
]
|
||||
completed_calls = [
|
||||
c for c in update_calls
|
||||
if c.kwargs.get("status") == "completed"
|
||||
]
|
||||
assert len(running_calls) >= 1, "Expected a 'running' status update"
|
||||
assert len(completed_calls) >= 1, "Expected a 'completed' status update"
|
||||
handler.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_job_basic_failed(self, cron_manager, mock_db):
|
||||
"""A basic job transitions to 'failed' and records the error."""
|
||||
mock_db.get_cron_job.return_value = CronJob(
|
||||
job_id="basic-fail",
|
||||
name="Fail",
|
||||
job_type="basic",
|
||||
enabled=True,
|
||||
cron_expression="0 9 * * *",
|
||||
)
|
||||
# Remove the handler so _run_basic_job raises RuntimeError
|
||||
if "basic-fail" in cron_manager._basic_handlers:
|
||||
del cron_manager._basic_handlers["basic-fail"]
|
||||
|
||||
await cron_manager._run_job("basic-fail")
|
||||
|
||||
# final update should have status="failed" and last_error set
|
||||
last_call = mock_db.update_cron_job.call_args_list[-1]
|
||||
assert last_call.kwargs["status"] == "failed"
|
||||
assert last_call.kwargs["last_error"] is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_job_run_once_deletes_after_completion(self, cron_manager, mock_db):
|
||||
"""A run_once job is deleted after it completes successfully."""
|
||||
job = CronJob(
|
||||
job_id="once-job",
|
||||
name="Once",
|
||||
job_type="basic",
|
||||
enabled=True,
|
||||
run_once=True,
|
||||
cron_expression="0 9 * * *",
|
||||
)
|
||||
mock_db.get_cron_job.return_value = job
|
||||
|
||||
handler = MagicMock(return_value=None)
|
||||
cron_manager._basic_handlers["once-job"] = handler
|
||||
|
||||
with patch.object(cron_manager, "delete_job") as mock_delete:
|
||||
await cron_manager._run_job("once-job")
|
||||
|
||||
mock_delete.assert_awaited_once_with("once-job")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_job_active_agent_raises_on_missing_session(self, cron_manager, mock_db):
|
||||
"""run_job on an active_agent job without session payload raises."""
|
||||
job = CronJob(
|
||||
job_id="aa-no-session",
|
||||
name="NoSession",
|
||||
job_type="active_agent",
|
||||
enabled=True,
|
||||
cron_expression="0 9 * * *",
|
||||
payload={}, # No "session" key
|
||||
)
|
||||
mock_db.get_cron_job.return_value = job
|
||||
|
||||
await cron_manager._run_job("aa-no-session")
|
||||
|
||||
# Should not crash the manager; error is caught and logged,
|
||||
# status should be "failed"
|
||||
last_call = mock_db.update_cron_job.call_args_list[-1]
|
||||
assert last_call.kwargs["status"] == "failed"
|
||||
assert "missing session" in (last_call.kwargs.get("last_error") or "").lower()
|
||||
|
||||
|
||||
# ---- update_job and sync_from_db ----
|
||||
|
||||
|
||||
class TestUpdateAndSyncEdgeCases:
|
||||
"""Edge cases for update_job and sync_from_db."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_job_reschedules_when_enabled(self, cron_manager, mock_db):
|
||||
"""update_job re-schedules a job that was previously disabled."""
|
||||
job_id = "re-enable-job"
|
||||
updated_job = CronJob(
|
||||
job_id=job_id,
|
||||
name="ReEnabled",
|
||||
job_type="basic",
|
||||
cron_expression="0 9 * * *",
|
||||
enabled=True, # Now enabled
|
||||
)
|
||||
mock_db.update_cron_job.return_value = updated_job
|
||||
|
||||
with patch.object(cron_manager, "_schedule_job") as mock_schedule:
|
||||
result = await cron_manager.update_job(job_id, enabled=True)
|
||||
|
||||
assert result == updated_job
|
||||
mock_schedule.assert_called_once_with(updated_job)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_job_removes_scheduled_when_disabled(self, cron_manager, mock_db):
|
||||
"""update_job removes the job from the scheduler when disabled."""
|
||||
job_id = "disable-job"
|
||||
updated_job = CronJob(
|
||||
job_id=job_id,
|
||||
name="Disabled",
|
||||
job_type="basic",
|
||||
cron_expression="0 9 * * *",
|
||||
enabled=False, # Now disabled, should not schedule
|
||||
)
|
||||
mock_db.update_cron_job.return_value = updated_job
|
||||
|
||||
with patch.object(cron_manager, "_remove_scheduled") as mock_remove:
|
||||
with patch.object(cron_manager, "_schedule_job") as mock_schedule:
|
||||
result = await cron_manager.update_job(job_id, enabled=False)
|
||||
|
||||
assert result == updated_job
|
||||
mock_remove.assert_called_once_with(job_id)
|
||||
mock_schedule.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_from_db_schedules_basic_with_handler(self, cron_manager, mock_db):
|
||||
"""sync_from_db schedules basic jobs when their handler is registered."""
|
||||
job = CronJob(
|
||||
job_id="sync-basic",
|
||||
name="SyncBasic",
|
||||
job_type="basic",
|
||||
cron_expression="0 9 * * *",
|
||||
enabled=True,
|
||||
persistent=True,
|
||||
)
|
||||
mock_db.list_cron_jobs.return_value = [job]
|
||||
cron_manager._basic_handlers["sync-basic"] = MagicMock()
|
||||
|
||||
with patch.object(cron_manager, "_schedule_job") as mock_schedule:
|
||||
await cron_manager.sync_from_db()
|
||||
|
||||
mock_schedule.assert_called_once_with(job)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_from_db_skips_basic_without_handler(self, cron_manager, mock_db):
|
||||
"""sync_from_db skips basic jobs that have no registered handler."""
|
||||
job = CronJob(
|
||||
job_id="orphan-basic",
|
||||
name="Orphan",
|
||||
job_type="basic",
|
||||
cron_expression="0 9 * * *",
|
||||
enabled=True,
|
||||
persistent=True,
|
||||
)
|
||||
mock_db.list_cron_jobs.return_value = [job]
|
||||
|
||||
with patch.object(cron_manager, "_schedule_job") as mock_schedule:
|
||||
await cron_manager.sync_from_db()
|
||||
|
||||
mock_schedule.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_from_db_schedules_active_agent_without_handler(self, cron_manager, mock_db):
|
||||
"""Active-agent jobs are scheduled regardless of handler registration."""
|
||||
job = CronJob(
|
||||
job_id="sync-active",
|
||||
name="SyncActive",
|
||||
job_type="active_agent",
|
||||
cron_expression="0 9 * * *",
|
||||
enabled=True,
|
||||
persistent=True,
|
||||
)
|
||||
mock_db.list_cron_jobs.return_value = [job]
|
||||
|
||||
with patch.object(cron_manager, "_schedule_job") as mock_schedule:
|
||||
await cron_manager.sync_from_db()
|
||||
|
||||
mock_schedule.assert_called_once_with(job)
|
||||
|
||||
|
||||
# ---- _schedule_job trigger variants ----
|
||||
|
||||
|
||||
class TestScheduleJobTriggers:
|
||||
"""Trigger type selection in _schedule_job."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_interval_trigger(self, cron_manager, mock_context):
|
||||
"""_schedule_job creates an IntervalTrigger when payload has interval_seconds."""
|
||||
job = CronJob(
|
||||
job_id="interval-trigger",
|
||||
name="Interval",
|
||||
job_type="basic",
|
||||
cron_expression="0 9 * * *",
|
||||
enabled=True,
|
||||
payload={"interval_seconds": 600},
|
||||
)
|
||||
mock_db = cron_manager.db
|
||||
mock_db.list_cron_jobs = AsyncMock(return_value=[])
|
||||
mock_db.update_cron_job = AsyncMock()
|
||||
await cron_manager.start(mock_context)
|
||||
|
||||
cron_manager._schedule_job(job)
|
||||
|
||||
aps_job = cron_manager.scheduler.get_job("interval-trigger")
|
||||
assert aps_job is not None
|
||||
# The trigger should be an IntervalTrigger
|
||||
from apscheduler.triggers.interval import IntervalTrigger
|
||||
|
||||
assert isinstance(aps_job.trigger, IntervalTrigger)
|
||||
assert aps_job.trigger.interval.total_seconds() == 600
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_invalid_cron_raises_scheduling_error(self, cron_manager, mock_context):
|
||||
"""An invalid cron expression raises CronJobSchedulingError."""
|
||||
job = CronJob(
|
||||
job_id="bad-cron",
|
||||
name="BadCron",
|
||||
job_type="basic",
|
||||
cron_expression="not-a-valid-cron",
|
||||
enabled=True,
|
||||
)
|
||||
mock_db = cron_manager.db
|
||||
mock_db.list_cron_jobs = AsyncMock(return_value=[])
|
||||
mock_db.update_cron_job = AsyncMock()
|
||||
await cron_manager.start(mock_context)
|
||||
|
||||
with pytest.raises(CronJobSchedulingError):
|
||||
cron_manager._schedule_job(job)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_run_once_without_run_at_raises(self, cron_manager, mock_context):
|
||||
"""A run_once job without run_at in payload or expression raises CronJobSchedulingError."""
|
||||
job = CronJob(
|
||||
job_id="no-run-at",
|
||||
name="NoRunAt",
|
||||
job_type="active_agent",
|
||||
cron_expression=None,
|
||||
enabled=True,
|
||||
run_once=True,
|
||||
payload={},
|
||||
)
|
||||
mock_db = cron_manager.db
|
||||
mock_db.list_cron_jobs = AsyncMock(return_value=[])
|
||||
mock_db.update_cron_job = AsyncMock()
|
||||
await cron_manager.start(mock_context)
|
||||
|
||||
with pytest.raises(CronJobSchedulingError):
|
||||
cron_manager._schedule_job(job)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schedule_auto_starts_when_not_started(self, cron_manager):
|
||||
"""_schedule_job auto-starts the scheduler if _started is False."""
|
||||
assert cron_manager._started is False
|
||||
|
||||
job = CronJob(
|
||||
job_id="auto-start",
|
||||
name="AutoStart",
|
||||
job_type="basic",
|
||||
cron_expression="0 9 * * *",
|
||||
enabled=True,
|
||||
)
|
||||
cron_manager.db.list_cron_jobs = AsyncMock(return_value=[])
|
||||
cron_manager.db.update_cron_job = AsyncMock()
|
||||
|
||||
cron_manager._schedule_job(job)
|
||||
|
||||
assert cron_manager._started is True
|
||||
assert cron_manager.scheduler.get_job("auto-start") is not None
|
||||
@@ -0,0 +1,254 @@
|
||||
"""Tests for BaseDatabase abstract interface and initialization.
|
||||
|
||||
Verifies the abstract method contract, engine creation, session factory
|
||||
setup, and the ``get_db`` async context manager.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
from abc import ABC
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.db import BaseDatabase
|
||||
|
||||
|
||||
class TestBaseDatabaseAbstract:
|
||||
"""Tests for BaseDatabase's abstract nature and method count."""
|
||||
|
||||
def test_is_abstract_class(self):
|
||||
"""BaseDatabase uses ABCMeta and cannot be directly instantiated."""
|
||||
assert isinstance(BaseDatabase, ABC)
|
||||
with pytest.raises(TypeError, match="abstract"):
|
||||
BaseDatabase()
|
||||
|
||||
def test_abstract_method_count(self):
|
||||
"""BaseDatabase defines a known number of abstract methods."""
|
||||
abs_methods = BaseDatabase.__abstractmethods__
|
||||
# The interface is large -- expect >= 90 abstract methods
|
||||
assert len(abs_methods) >= 90
|
||||
# Key abstract methods should be present
|
||||
assert "create_cron_job" in abs_methods
|
||||
assert "update_cron_job" in abs_methods
|
||||
assert "delete_cron_job" in abs_methods
|
||||
assert "get_cron_job" in abs_methods
|
||||
assert "list_cron_jobs" in abs_methods
|
||||
assert "get_conversations" in abs_methods
|
||||
assert "create_conversation" in abs_methods
|
||||
assert "insert_platform_stats" in abs_methods
|
||||
assert "insert_persona" in abs_methods
|
||||
assert "insert_api_key" in abs_methods # kept as create_api_key is the actual name
|
||||
# Deprecated methods are still abstract
|
||||
assert "get_base_stats" in abs_methods
|
||||
assert "get_total_message_count" in abs_methods
|
||||
|
||||
def test_abstract_methods_have_docstrings_or_impl(self):
|
||||
"""All abstract methods should have ... or docstrings (no syntax errors at module level)."""
|
||||
abs_methods = BaseDatabase.__abstractmethods__
|
||||
for name in abs_methods:
|
||||
method = getattr(BaseDatabase, name)
|
||||
# Should not raise AttributeError or similar
|
||||
assert callable(method)
|
||||
|
||||
def test_initialize_is_not_abstract(self):
|
||||
"""initialize() has a concrete default implementation."""
|
||||
assert "initialize" not in BaseDatabase.__abstractmethods__
|
||||
|
||||
def test_get_db_is_not_abstract(self):
|
||||
"""get_db() has a concrete implementation (async context manager)."""
|
||||
assert "get_db" not in BaseDatabase.__abstractmethods__
|
||||
|
||||
|
||||
class TestBaseDatabaseInit:
|
||||
"""Tests for BaseDatabase.__init__ and engine setup."""
|
||||
|
||||
def test_init_sets_inited_false(self):
|
||||
"""inited starts as False."""
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine") as mock_engine,
|
||||
patch("astrbot.core.db.async_sessionmaker") as mock_smaker,
|
||||
):
|
||||
mock_engine.return_value = MagicMock()
|
||||
mock_smaker.return_value = MagicMock()
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///test.db"
|
||||
db = BaseDatabase()
|
||||
|
||||
assert db.inited is False
|
||||
|
||||
def test_init_sqlite_adds_connect_args_timeout(self):
|
||||
"""SQLite URL adds timeout=30 to connect_args."""
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine") as mock_engine,
|
||||
patch("astrbot.core.db.async_sessionmaker") as mock_smaker,
|
||||
):
|
||||
mock_engine.return_value = MagicMock()
|
||||
mock_smaker.return_value = MagicMock()
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///test.db"
|
||||
db = BaseDatabase() # noqa: F841
|
||||
|
||||
mock_engine.assert_called_once()
|
||||
call_kwargs = mock_engine.call_args.kwargs
|
||||
assert call_kwargs["connect_args"] == {"timeout": 30}
|
||||
|
||||
def test_init_non_sqlite_does_not_add_connect_args(self):
|
||||
"""Non-SQLite URL omits connect_args (no timeout)."""
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine") as mock_engine,
|
||||
patch("astrbot.core.db.async_sessionmaker") as mock_smaker,
|
||||
):
|
||||
mock_engine.return_value = MagicMock()
|
||||
mock_smaker.return_value = MagicMock()
|
||||
BaseDatabase.DATABASE_URL = "postgresql+asyncpg://localhost/db"
|
||||
db = BaseDatabase() # noqa: F841
|
||||
|
||||
call_kwargs = mock_engine.call_args.kwargs
|
||||
assert "connect_args" not in call_kwargs or call_kwargs["connect_args"] == {}
|
||||
|
||||
def test_init_creates_async_session_local(self):
|
||||
"""A session factory is created from the engine."""
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine") as mock_engine,
|
||||
patch("astrbot.core.db.async_sessionmaker") as mock_smaker,
|
||||
):
|
||||
mock_engine.return_value = MagicMock()
|
||||
mock_smaker.return_value = MagicMock()
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///test.db"
|
||||
db = BaseDatabase() # noqa: F841
|
||||
|
||||
mock_smaker.assert_called_once()
|
||||
# The session maker is bound to the engine with expire_on_commit=False
|
||||
assert mock_smaker.call_args.kwargs.get("expire_on_commit") is False
|
||||
|
||||
def test_init_engine_echo_false_future_true(self):
|
||||
"""Engine is created with echo=False and future=True."""
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine") as mock_engine,
|
||||
patch("astrbot.core.db.async_sessionmaker") as mock_smaker,
|
||||
):
|
||||
mock_engine.return_value = MagicMock()
|
||||
mock_smaker.return_value = MagicMock()
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///test.db"
|
||||
db = BaseDatabase() # noqa: F841
|
||||
|
||||
call_kwargs = mock_engine.call_args.kwargs
|
||||
assert call_kwargs["echo"] is False
|
||||
assert call_kwargs["future"] is True
|
||||
|
||||
def test_same_url_passed_to_engine(self):
|
||||
"""DATABASE_URL is passed to create_async_engine."""
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine") as mock_engine,
|
||||
patch("astrbot.core.db.async_sessionmaker") as mock_smaker,
|
||||
):
|
||||
mock_engine.return_value = MagicMock()
|
||||
mock_smaker.return_value = MagicMock()
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///custom.db"
|
||||
db = BaseDatabase() # noqa: F841
|
||||
|
||||
assert mock_engine.call_args[0][0] == "sqlite+aiosqlite:///custom.db"
|
||||
|
||||
|
||||
class TestBaseDatabaseInitialize:
|
||||
"""Tests for the concrete initialize method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_does_not_raise(self):
|
||||
"""initialize() is concrete and can be called without error."""
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine"),
|
||||
patch("astrbot.core.db.async_sessionmaker"),
|
||||
):
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///test.db"
|
||||
db = BaseDatabase()
|
||||
|
||||
await db.initialize() # should not raise
|
||||
|
||||
|
||||
class TestBaseDatabaseGetDb:
|
||||
"""Tests for the get_db async context manager."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_db_calls_initialize_when_not_inited(self):
|
||||
"""get_db calls initialize() if inited is False."""
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine"),
|
||||
patch("astrbot.core.db.async_sessionmaker"),
|
||||
):
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///test.db"
|
||||
db = BaseDatabase()
|
||||
|
||||
db.inited = False
|
||||
|
||||
# Mock the session factory
|
||||
mock_session = MagicMock(spec=AsyncSession)
|
||||
mock_cm = AsyncMock()
|
||||
mock_cm.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
db.AsyncSessionLocal = MagicMock(return_value=mock_cm)
|
||||
|
||||
with patch.object(db, "initialize", new_callable=AsyncMock) as mock_init:
|
||||
async with db.get_db() as session:
|
||||
assert session == mock_session
|
||||
|
||||
mock_init.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_db_does_not_call_initialize_when_inited(self):
|
||||
"""get_db skips initialize() if inited is already True."""
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine"),
|
||||
patch("astrbot.core.db.async_sessionmaker"),
|
||||
):
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///test.db"
|
||||
db = BaseDatabase()
|
||||
|
||||
db.inited = True
|
||||
|
||||
mock_session = MagicMock(spec=AsyncSession)
|
||||
mock_cm = AsyncMock()
|
||||
mock_cm.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
db.AsyncSessionLocal = MagicMock(return_value=mock_cm)
|
||||
|
||||
with patch.object(db, "initialize", new_callable=AsyncMock) as mock_init:
|
||||
async with db.get_db() as session:
|
||||
assert session == mock_session
|
||||
|
||||
mock_init.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_db_sets_inited_after_initialize(self):
|
||||
"""inited is set to True after initialize completes."""
|
||||
with (
|
||||
patch.object(BaseDatabase, "__abstractmethods__", frozenset()),
|
||||
patch("astrbot.core.db.create_async_engine"),
|
||||
patch("astrbot.core.db.async_sessionmaker"),
|
||||
):
|
||||
BaseDatabase.DATABASE_URL = "sqlite+aiosqlite:///test.db"
|
||||
db = BaseDatabase()
|
||||
|
||||
db.inited = False
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_cm = AsyncMock()
|
||||
mock_cm.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_cm.__aexit__ = AsyncMock(return_value=None)
|
||||
db.AsyncSessionLocal = MagicMock(return_value=mock_cm)
|
||||
|
||||
async with db.get_db():
|
||||
pass
|
||||
|
||||
assert db.inited is True
|
||||
@@ -0,0 +1,454 @@
|
||||
"""Tests for database PO (Persistent Object) model classes.
|
||||
|
||||
These tests verify construction, field defaults, and type correctness
|
||||
for all model classes defined in ``astrbot.core.db.po``.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.db.po import (
|
||||
ApiKey,
|
||||
Attachment,
|
||||
ChatUIProject,
|
||||
CommandConfig,
|
||||
CommandConflict,
|
||||
Conversation,
|
||||
ConversationV2,
|
||||
CronJob,
|
||||
Persona,
|
||||
PersonaFolder,
|
||||
Personality,
|
||||
PlatformMessageHistory,
|
||||
PlatformSession,
|
||||
PlatformStat,
|
||||
Preference,
|
||||
ProviderStat,
|
||||
SessionProjectRelation,
|
||||
Stats,
|
||||
TimestampMixin,
|
||||
WebChatThread,
|
||||
)
|
||||
|
||||
|
||||
class TestCronJob:
|
||||
"""Tests for the CronJob SQLModel."""
|
||||
|
||||
def test_minimal_construction(self):
|
||||
"""CronJob can be created with only required fields."""
|
||||
job = CronJob(name="test-job", job_type="basic")
|
||||
assert job.name == "test-job"
|
||||
assert job.job_type == "basic"
|
||||
assert job.enabled is True
|
||||
assert job.persistent is True
|
||||
assert job.run_once is False
|
||||
assert job.status == "scheduled"
|
||||
assert job.cron_expression is None
|
||||
assert job.job_id is not None # auto-generated via default_factory
|
||||
|
||||
def test_full_construction(self):
|
||||
"""CronJob accepts all optional fields."""
|
||||
job = CronJob(
|
||||
name="full-job",
|
||||
job_type="active_agent",
|
||||
cron_expression="0 9 * * *",
|
||||
timezone="UTC",
|
||||
payload={"session": "test:group:1"},
|
||||
description="A full job",
|
||||
enabled=False,
|
||||
persistent=True,
|
||||
run_once=True,
|
||||
status="running",
|
||||
)
|
||||
assert job.name == "full-job"
|
||||
assert job.job_type == "active_agent"
|
||||
assert job.cron_expression == "0 9 * * *"
|
||||
assert job.timezone == "UTC"
|
||||
assert job.payload == {"session": "test:group:1"}
|
||||
assert job.description == "A full job"
|
||||
assert job.enabled is False
|
||||
assert job.run_once is True
|
||||
assert job.status == "running"
|
||||
|
||||
def test_auto_timestamps(self):
|
||||
"""CronJob inherits TimestampMixin which auto-generates created_at/updated_at."""
|
||||
job = CronJob(name="ts-test", job_type="basic")
|
||||
assert isinstance(job.created_at, datetime)
|
||||
assert isinstance(job.updated_at, datetime)
|
||||
|
||||
def test_job_id_auto_generated(self):
|
||||
"""Each instance gets a unique job_id."""
|
||||
job1 = CronJob(name="a", job_type="basic")
|
||||
job2 = CronJob(name="b", job_type="basic")
|
||||
assert job1.job_id != job2.job_id
|
||||
|
||||
|
||||
class TestConversationV2:
|
||||
"""Tests for the ConversationV2 SQLModel."""
|
||||
|
||||
def test_minimal_construction(self):
|
||||
"""ConversationV2 can be created with required fields."""
|
||||
conv = ConversationV2(platform_id="qq", user_id="user-1")
|
||||
assert conv.platform_id == "qq"
|
||||
assert conv.user_id == "user-1"
|
||||
assert conv.conversation_id is not None
|
||||
assert conv.title is None
|
||||
assert conv.persona_id is None
|
||||
assert conv.token_usage == 0
|
||||
|
||||
def test_full_construction(self):
|
||||
"""ConversationV2 accepts all optional fields."""
|
||||
conv = ConversationV2(
|
||||
platform_id="webchat",
|
||||
user_id="admin",
|
||||
content=[{"role": "user", "content": "hi"}],
|
||||
title="Chat",
|
||||
persona_id="p1",
|
||||
token_usage=42,
|
||||
)
|
||||
assert conv.content == [{"role": "user", "content": "hi"}]
|
||||
assert conv.title == "Chat"
|
||||
assert conv.persona_id == "p1"
|
||||
assert conv.token_usage == 42
|
||||
|
||||
|
||||
class TestConversationDataclass:
|
||||
"""Tests for the deprecated Conversation dataclass."""
|
||||
|
||||
def test_construction(self):
|
||||
"""Conversation dataclass sets all fields."""
|
||||
conv = Conversation(
|
||||
platform_id="qq",
|
||||
user_id="u1",
|
||||
cid="abc-123",
|
||||
history="[{}]",
|
||||
title="My Chat",
|
||||
persona_id="p1",
|
||||
created_at=1000,
|
||||
updated_at=2000,
|
||||
token_usage=50,
|
||||
)
|
||||
assert conv.platform_id == "qq"
|
||||
assert conv.user_id == "u1"
|
||||
assert conv.cid == "abc-123"
|
||||
assert conv.history == "[{}]"
|
||||
assert conv.title == "My Chat"
|
||||
assert conv.persona_id == "p1"
|
||||
assert conv.created_at == 1000
|
||||
assert conv.updated_at == 2000
|
||||
assert conv.token_usage == 50
|
||||
|
||||
def test_defaults(self):
|
||||
"""Conversation dataclass has sensible defaults."""
|
||||
conv = Conversation(platform_id="qq", user_id="u1", cid="x")
|
||||
assert conv.history == ""
|
||||
assert conv.title == ""
|
||||
assert conv.persona_id == ""
|
||||
assert conv.created_at == 0
|
||||
assert conv.token_usage == 0
|
||||
|
||||
|
||||
class TestPersona:
|
||||
"""Tests for the Persona SQLModel."""
|
||||
|
||||
def test_minimal_construction(self):
|
||||
"""Persona with only required fields."""
|
||||
p = Persona(persona_id="helper", system_prompt="You are helpful.")
|
||||
assert p.persona_id == "helper"
|
||||
assert p.system_prompt == "You are helpful."
|
||||
assert p.begin_dialogs is None
|
||||
assert p.tools is None
|
||||
assert p.skills is None
|
||||
assert p.custom_error_message is None
|
||||
assert p.sort_order == 0
|
||||
|
||||
def test_with_all_fields(self):
|
||||
"""Persona with all optional fields."""
|
||||
p = Persona(
|
||||
persona_id="custom",
|
||||
system_prompt="Be concise.",
|
||||
begin_dialogs=["Hello"],
|
||||
tools=["search", "calc"],
|
||||
skills=["weather"],
|
||||
custom_error_message="Sorry, try later.",
|
||||
folder_id="folder-1",
|
||||
sort_order=5,
|
||||
)
|
||||
assert p.begin_dialogs == ["Hello"]
|
||||
assert p.tools == ["search", "calc"]
|
||||
assert p.skills == ["weather"]
|
||||
assert p.custom_error_message == "Sorry, try later."
|
||||
assert p.folder_id == "folder-1"
|
||||
assert p.sort_order == 5
|
||||
|
||||
|
||||
class TestPersonaFolder:
|
||||
"""Tests for the PersonaFolder SQLModel."""
|
||||
|
||||
def test_construction(self):
|
||||
"""PersonaFolder creation with required fields."""
|
||||
folder = PersonaFolder(name="My Folder")
|
||||
assert folder.name == "My Folder"
|
||||
assert folder.folder_id is not None
|
||||
assert folder.parent_id is None
|
||||
assert folder.description is None
|
||||
assert folder.sort_order == 0
|
||||
|
||||
def test_nested_folder(self):
|
||||
"""PersonaFolder can have a parent_id."""
|
||||
child = PersonaFolder(name="Child", parent_id="parent-uuid", description="Nested")
|
||||
assert child.parent_id == "parent-uuid"
|
||||
assert child.description == "Nested"
|
||||
|
||||
|
||||
class TestApiKey:
|
||||
"""Tests for the ApiKey SQLModel."""
|
||||
|
||||
def test_construction(self):
|
||||
"""ApiKey with required fields."""
|
||||
key = ApiKey(
|
||||
name="dev-key",
|
||||
key_hash="sha256:abc123",
|
||||
key_prefix="astr_",
|
||||
created_by="admin",
|
||||
)
|
||||
assert key.name == "dev-key"
|
||||
assert key.key_hash == "sha256:abc123"
|
||||
assert key.key_prefix == "astr_"
|
||||
assert key.created_by == "admin"
|
||||
assert key.scopes is None
|
||||
assert key.last_used_at is None
|
||||
assert key.expires_at is None
|
||||
assert key.revoked_at is None
|
||||
assert key.key_id is not None
|
||||
|
||||
def test_with_scopes(self):
|
||||
"""ApiKey can have scopes and expiry."""
|
||||
expires = datetime.now(timezone.utc)
|
||||
key = ApiKey(
|
||||
name="scoped-key",
|
||||
key_hash="sha256:xyz",
|
||||
key_prefix="astr_",
|
||||
created_by="admin",
|
||||
scopes=["read", "write"],
|
||||
expires_at=expires,
|
||||
)
|
||||
assert key.scopes == ["read", "write"]
|
||||
assert key.expires_at == expires
|
||||
|
||||
|
||||
class TestPlatformStat:
|
||||
"""Tests for the PlatformStat SQLModel."""
|
||||
|
||||
def test_construction(self):
|
||||
"""PlatformStat with all fields."""
|
||||
ts = datetime.now(timezone.utc)
|
||||
stat = PlatformStat(
|
||||
timestamp=ts,
|
||||
platform_id="qq_bot",
|
||||
platform_type="aiocqhttp",
|
||||
count=5,
|
||||
)
|
||||
assert stat.timestamp == ts
|
||||
assert stat.platform_id == "qq_bot"
|
||||
assert stat.platform_type == "aiocqhttp"
|
||||
assert stat.count == 5
|
||||
|
||||
def test_default_count(self):
|
||||
"""PlatformStat defaults count to 0."""
|
||||
stat = PlatformStat(
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
platform_id="test",
|
||||
platform_type="test",
|
||||
)
|
||||
assert stat.count == 0
|
||||
|
||||
|
||||
class TestProviderStat:
|
||||
"""Tests for the ProviderStat SQLModel."""
|
||||
|
||||
def test_construction(self):
|
||||
"""ProviderStat with required fields."""
|
||||
ps = ProviderStat(umo="test:private:1", provider_id="openai")
|
||||
assert ps.umo == "test:private:1"
|
||||
assert ps.provider_id == "openai"
|
||||
assert ps.agent_type == "internal"
|
||||
assert ps.status == "completed"
|
||||
assert ps.token_input_other == 0
|
||||
assert ps.token_input_cached == 0
|
||||
assert ps.token_output == 0
|
||||
|
||||
def test_with_stats(self):
|
||||
"""ProviderStat records token and timing stats."""
|
||||
ps = ProviderStat(
|
||||
umo="test:private:1",
|
||||
provider_id="anthropic",
|
||||
provider_model="claude-3",
|
||||
conversation_id="conv-1",
|
||||
status="completed",
|
||||
agent_type="cron",
|
||||
token_input_other=100,
|
||||
token_input_cached=50,
|
||||
token_output=200,
|
||||
start_time=1000.0,
|
||||
end_time=1005.0,
|
||||
time_to_first_token=0.5,
|
||||
)
|
||||
assert ps.provider_model == "claude-3"
|
||||
assert ps.conversation_id == "conv-1"
|
||||
assert ps.token_input_other == 100
|
||||
assert ps.token_input_cached == 50
|
||||
assert ps.token_output == 200
|
||||
assert ps.start_time == 1000.0
|
||||
assert ps.end_time == 1005.0
|
||||
assert ps.time_to_first_token == 0.5
|
||||
|
||||
|
||||
class TestPreference:
|
||||
"""Tests for the Preference SQLModel."""
|
||||
|
||||
def test_construction(self):
|
||||
"""Preference with required fields."""
|
||||
pref = Preference(
|
||||
scope="plugin",
|
||||
scope_id="my_plugin",
|
||||
key="theme",
|
||||
value={"color": "dark"},
|
||||
)
|
||||
assert pref.scope == "plugin"
|
||||
assert pref.scope_id == "my_plugin"
|
||||
assert pref.key == "theme"
|
||||
assert pref.value == {"color": "dark"}
|
||||
|
||||
|
||||
class TestOtherModels:
|
||||
"""Tests for remaining SQLModel classes."""
|
||||
|
||||
def test_attachment(self):
|
||||
"""Attachment model construction."""
|
||||
att = Attachment(path="/tmp/file.png", type="image", mime_type="image/png")
|
||||
assert att.path == "/tmp/file.png"
|
||||
assert att.type == "image"
|
||||
assert att.mime_type == "image/png"
|
||||
assert att.attachment_id is not None
|
||||
|
||||
def test_webchat_thread(self):
|
||||
"""WebChatThread model construction."""
|
||||
thread = WebChatThread(
|
||||
creator="user-1",
|
||||
parent_session_id="session-1",
|
||||
parent_message_id=42,
|
||||
base_checkpoint_id="ckpt-1",
|
||||
selected_text="selected text",
|
||||
)
|
||||
assert thread.creator == "user-1"
|
||||
assert thread.parent_session_id == "session-1"
|
||||
assert thread.parent_message_id == 42
|
||||
assert thread.base_checkpoint_id == "ckpt-1"
|
||||
assert thread.selected_text == "selected text"
|
||||
assert thread.thread_id is not None
|
||||
|
||||
def test_platform_session(self):
|
||||
"""PlatformSession model construction."""
|
||||
sess = PlatformSession(
|
||||
creator="user-1",
|
||||
platform_id="webchat",
|
||||
display_name="My Chat",
|
||||
)
|
||||
assert sess.creator == "user-1"
|
||||
assert sess.platform_id == "webchat"
|
||||
assert sess.display_name == "My Chat"
|
||||
assert sess.is_group == 0
|
||||
assert sess.session_id is not None
|
||||
|
||||
def test_chatui_project(self):
|
||||
"""ChatUIProject model construction."""
|
||||
proj = ChatUIProject(
|
||||
creator="admin",
|
||||
title="My Project",
|
||||
emoji="star",
|
||||
description="A test project",
|
||||
)
|
||||
assert proj.creator == "admin"
|
||||
assert proj.title == "My Project"
|
||||
assert proj.emoji == "star"
|
||||
assert proj.description == "A test project"
|
||||
assert proj.project_id is not None
|
||||
|
||||
def test_session_project_relation(self):
|
||||
"""SessionProjectRelation model construction."""
|
||||
rel = SessionProjectRelation(session_id="sess-1", project_id="proj-1")
|
||||
assert rel.session_id == "sess-1"
|
||||
assert rel.project_id == "proj-1"
|
||||
|
||||
def test_command_config(self):
|
||||
"""CommandConfig model construction."""
|
||||
cc = CommandConfig(
|
||||
handler_full_name="plugin.command",
|
||||
plugin_name="test-plugin",
|
||||
module_path="plugins.test",
|
||||
original_command="/test",
|
||||
)
|
||||
assert cc.handler_full_name == "plugin.command"
|
||||
assert cc.plugin_name == "test-plugin"
|
||||
assert cc.module_path == "plugins.test"
|
||||
assert cc.original_command == "/test"
|
||||
assert cc.enabled is True
|
||||
assert cc.auto_managed is False
|
||||
|
||||
def test_command_conflict(self):
|
||||
"""CommandConflict model construction."""
|
||||
cf = CommandConflict(
|
||||
conflict_key="/greet",
|
||||
handler_full_name="p1.greet",
|
||||
plugin_name="plugin1",
|
||||
)
|
||||
assert cf.conflict_key == "/greet"
|
||||
assert cf.handler_full_name == "p1.greet"
|
||||
assert cf.plugin_name == "plugin1"
|
||||
assert cf.status == "pending"
|
||||
assert cf.auto_generated is False
|
||||
|
||||
def test_platform_message_history(self):
|
||||
"""PlatformMessageHistory model construction."""
|
||||
pmh = PlatformMessageHistory(
|
||||
platform_id="qq",
|
||||
user_id="user-1",
|
||||
content={"text": "hello"},
|
||||
)
|
||||
assert pmh.platform_id == "qq"
|
||||
assert pmh.user_id == "user-1"
|
||||
assert pmh.content == {"text": "hello"}
|
||||
assert pmh.sender_id is None
|
||||
assert pmh.sender_name is None
|
||||
assert pmh.llm_checkpoint_id is None
|
||||
|
||||
def test_timestamp_mixin_fields(self):
|
||||
"""TimestampMixin provides created_at and updated_at."""
|
||||
ts = datetime.now(timezone.utc)
|
||||
mixin = TimestampMixin()
|
||||
# created_at and updated_at have default factories
|
||||
assert isinstance(mixin.created_at, datetime)
|
||||
assert isinstance(mixin.updated_at, datetime)
|
||||
|
||||
def test_personality_typeddict(self):
|
||||
"""Personality TypedDict can be constructed with all keys."""
|
||||
personality: Personality = {
|
||||
"prompt": "You are a bot.",
|
||||
"name": "Bot",
|
||||
"begin_dialogs": ["Hello"],
|
||||
"mood_imitation_dialogs": ["Hi there"],
|
||||
"tools": ["search"],
|
||||
"skills": ["weather"],
|
||||
"custom_error_message": "Oops",
|
||||
"_begin_dialogs_processed": [{"role": "user", "content": "Hello"}],
|
||||
"_mood_imitation_dialogs_processed": "Hi there",
|
||||
}
|
||||
assert personality["prompt"] == "You are a bot."
|
||||
assert personality["name"] == "Bot"
|
||||
|
||||
def test_stats_dataclass(self):
|
||||
"""Stats dataclass holds a list of Platforms."""
|
||||
stats = Stats()
|
||||
assert stats.platform == []
|
||||
@@ -0,0 +1,211 @@
|
||||
"""Mock-based unit tests for default config constants and structure."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.config.default import (
|
||||
DB_PATH,
|
||||
DEFAULT_CONFIG,
|
||||
DEFAULT_VALUE_MAP,
|
||||
VERSION,
|
||||
WEBHOOK_SUPPORTED_PLATFORMS,
|
||||
)
|
||||
|
||||
|
||||
class TestVersionConstant:
|
||||
"""Tests for the VERSION constant."""
|
||||
|
||||
def test_version_is_string(self):
|
||||
assert isinstance(VERSION, str)
|
||||
|
||||
def test_version_format(self):
|
||||
parts = VERSION.split(".")
|
||||
assert len(parts) == 3
|
||||
for p in parts:
|
||||
assert p.isdigit()
|
||||
|
||||
|
||||
class TestDBPath:
|
||||
"""Tests for DB_PATH."""
|
||||
|
||||
@patch("astrbot.core.config.default.get_astrbot_data_path", return_value="/data")
|
||||
def test_db_path_uses_data_path(self, mock_get_path):
|
||||
"""DB_PATH should join data path with the database filename."""
|
||||
# Reimport to pick up the patched value
|
||||
import importlib
|
||||
|
||||
from astrbot.core.config import default as default_mod
|
||||
|
||||
importlib.reload(default_mod)
|
||||
assert "data_v4.db" in default_mod.DB_PATH
|
||||
|
||||
def test_db_path_ends_with_db(self):
|
||||
assert DB_PATH.endswith(".db")
|
||||
|
||||
|
||||
class TestDEFAULT_VALUE_MAP:
|
||||
"""Tests for DEFAULT_VALUE_MAP structure."""
|
||||
|
||||
def test_contains_expected_types(self):
|
||||
expected = {"int", "float", "bool", "string", "text", "list", "file", "object", "template_list"}
|
||||
assert set(DEFAULT_VALUE_MAP.keys()) == expected
|
||||
|
||||
def test_default_values_are_correct_types(self):
|
||||
assert isinstance(DEFAULT_VALUE_MAP["int"], int)
|
||||
assert isinstance(DEFAULT_VALUE_MAP["float"], float)
|
||||
assert isinstance(DEFAULT_VALUE_MAP["bool"], bool)
|
||||
assert isinstance(DEFAULT_VALUE_MAP["string"], str)
|
||||
assert isinstance(DEFAULT_VALUE_MAP["text"], str)
|
||||
assert isinstance(DEFAULT_VALUE_MAP["list"], list)
|
||||
assert isinstance(DEFAULT_VALUE_MAP["file"], list)
|
||||
assert isinstance(DEFAULT_VALUE_MAP["object"], dict)
|
||||
assert isinstance(DEFAULT_VALUE_MAP["template_list"], list)
|
||||
|
||||
def test_specific_default_values(self):
|
||||
assert DEFAULT_VALUE_MAP["int"] == 0
|
||||
assert DEFAULT_VALUE_MAP["float"] == 0.0
|
||||
assert DEFAULT_VALUE_MAP["bool"] is False
|
||||
assert DEFAULT_VALUE_MAP["string"] == ""
|
||||
assert DEFAULT_VALUE_MAP["list"] == []
|
||||
assert DEFAULT_VALUE_MAP["object"] == {}
|
||||
assert DEFAULT_VALUE_MAP["template_list"] == []
|
||||
|
||||
|
||||
class TestWEBHOOK_SUPPORTED_PLATFORMS:
|
||||
"""Tests for the webhook platforms list."""
|
||||
|
||||
def test_is_list_of_strings(self):
|
||||
assert isinstance(WEBHOOK_SUPPORTED_PLATFORMS, list)
|
||||
for p in WEBHOOK_SUPPORTED_PLATFORMS:
|
||||
assert isinstance(p, str)
|
||||
|
||||
def test_contains_key_platforms(self):
|
||||
assert "qq_official_webhook" in WEBHOOK_SUPPORTED_PLATFORMS
|
||||
assert "weixin_official_account" in WEBHOOK_SUPPORTED_PLATFORMS
|
||||
assert "slack" in WEBHOOK_SUPPORTED_PLATFORMS
|
||||
assert "lark" in WEBHOOK_SUPPORTED_PLATFORMS
|
||||
|
||||
|
||||
class TestDEFAULT_CONFIGStructure:
|
||||
"""Tests for the top-level keys and nested structure of DEFAULT_CONFIG."""
|
||||
|
||||
def test_top_level_keys(self):
|
||||
expected_keys = {
|
||||
"config_version",
|
||||
"platform_settings",
|
||||
"provider_sources",
|
||||
"provider",
|
||||
"provider_settings",
|
||||
"subagent_orchestrator",
|
||||
"provider_stt_settings",
|
||||
"provider_tts_settings",
|
||||
"provider_ltm_settings",
|
||||
"content_safety",
|
||||
"admins_id",
|
||||
"t2i",
|
||||
"http_proxy",
|
||||
"no_proxy",
|
||||
"dashboard",
|
||||
"platform",
|
||||
"platform_specific",
|
||||
"wake_prefix",
|
||||
"log_level",
|
||||
"persona",
|
||||
"timezone",
|
||||
"callback_api_base",
|
||||
"default_kb_collection",
|
||||
"plugin_set",
|
||||
"kb_names",
|
||||
"kb_fusion_top_k",
|
||||
"kb_final_top_k",
|
||||
"kb_agentic_mode",
|
||||
"disable_builtin_commands",
|
||||
"t2i_word_threshold",
|
||||
"t2i_strategy",
|
||||
"t2i_endpoint",
|
||||
"t2i_use_file_service",
|
||||
"t2i_active_template",
|
||||
"log_file_enable",
|
||||
"log_file_path",
|
||||
"log_file_max_mb",
|
||||
"temp_dir_max_size",
|
||||
"trace_enable",
|
||||
"trace_log_enable",
|
||||
"trace_log_path",
|
||||
"trace_log_max_mb",
|
||||
"pip_install_arg",
|
||||
"pypi_index_url",
|
||||
}
|
||||
for key in expected_keys:
|
||||
assert key in DEFAULT_CONFIG, f"Missing top-level key: {key}"
|
||||
|
||||
def test_config_version_is_two(self):
|
||||
assert DEFAULT_CONFIG["config_version"] == 2
|
||||
|
||||
def test_platform_settings_structure(self):
|
||||
ps = DEFAULT_CONFIG["platform_settings"]
|
||||
assert "unique_session" in ps
|
||||
assert "rate_limit" in ps
|
||||
assert "reply_prefix" in ps
|
||||
assert isinstance(ps["rate_limit"], dict)
|
||||
assert ps["rate_limit"]["strategy"] in ("stall", "discard")
|
||||
|
||||
def test_provider_ltm_settings_structure(self):
|
||||
ltm = DEFAULT_CONFIG["provider_ltm_settings"]
|
||||
assert "group_icl_enable" in ltm
|
||||
assert "group_message_max_cnt" in ltm
|
||||
assert "image_caption" in ltm
|
||||
assert "active_reply" in ltm
|
||||
assert isinstance(ltm["active_reply"], dict)
|
||||
assert ltm["active_reply"]["method"] == "possibility_reply"
|
||||
|
||||
def test_dashboard_settings(self):
|
||||
dash = DEFAULT_CONFIG["dashboard"]
|
||||
assert dash["enable"] is True
|
||||
assert dash["username"] == "astrbot"
|
||||
assert dash["port"] == 6185
|
||||
assert "ssl" in dash
|
||||
|
||||
def test_provider_settings_defaults(self):
|
||||
ps = DEFAULT_CONFIG["provider_settings"]
|
||||
assert ps["enable"] is True
|
||||
assert ps["default_provider_id"] == ""
|
||||
assert ps["agent_runner_type"] == "local"
|
||||
assert ps["llm_safety_mode"] is True
|
||||
|
||||
def test_admins_id_default(self):
|
||||
assert DEFAULT_CONFIG["admins_id"] == ["astrbot"]
|
||||
|
||||
def test_wake_prefix_default(self):
|
||||
assert DEFAULT_CONFIG["wake_prefix"] == ["/"]
|
||||
|
||||
def test_log_level_default(self):
|
||||
assert DEFAULT_CONFIG["log_level"] == "INFO"
|
||||
|
||||
def test_platform_specific_contains_lark_telegram_discord(self):
|
||||
ps = DEFAULT_CONFIG["platform_specific"]
|
||||
assert "lark" in ps
|
||||
assert "telegram" in ps
|
||||
assert "discord" in ps
|
||||
|
||||
def test_no_proxy_contains_localhost(self):
|
||||
assert "localhost" in DEFAULT_CONFIG["no_proxy"]
|
||||
|
||||
def test_subagent_orchestrator_defaults(self):
|
||||
sa = DEFAULT_CONFIG["subagent_orchestrator"]
|
||||
assert sa["main_enable"] is False
|
||||
assert isinstance(sa["agents"], list)
|
||||
|
||||
def test_quoted_message_parser_defaults(self):
|
||||
qmp = DEFAULT_CONFIG["provider_settings"]["quoted_message_parser"]
|
||||
assert qmp["max_component_chain_depth"] == 4
|
||||
assert qmp["max_forward_node_depth"] == 6
|
||||
assert qmp["max_forward_fetch"] == 32
|
||||
|
||||
def test_sandbox_defaults(self):
|
||||
sb = DEFAULT_CONFIG["provider_settings"]["sandbox"]
|
||||
assert sb["booter"] == "shipyard_neo"
|
||||
assert sb["shipyard_neo_ttl"] == 3600
|
||||
@@ -0,0 +1,396 @@
|
||||
"""Supplementary edge-case tests for EventBus.
|
||||
|
||||
Covers exception resilience in the dispatch loop, edge values in
|
||||
_print_event, empty pipeline mappings, and rapid event delivery.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from contextlib import suppress
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.event_bus import EventBus
|
||||
|
||||
|
||||
# ---- Fixtures ----
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def event_queue():
|
||||
return asyncio.Queue()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_scheduler():
|
||||
scheduler = MagicMock()
|
||||
scheduler.execute = AsyncMock()
|
||||
return scheduler
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_config_manager():
|
||||
config_mgr = MagicMock()
|
||||
config_mgr.get_conf_info = MagicMock(
|
||||
return_value={"id": "test-conf-id", "name": "Test Config"}
|
||||
)
|
||||
return config_mgr
|
||||
|
||||
|
||||
# ---- Exception resilience ----
|
||||
|
||||
|
||||
class TestDispatchExceptionResilience:
|
||||
"""The dispatch loop must survive scheduler exceptions."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exception_in_scheduler_does_not_crash_loop(
|
||||
self, event_queue, mock_config_manager
|
||||
):
|
||||
"""After a scheduler that raises, subsequent events are still processed."""
|
||||
processed_second = asyncio.Event()
|
||||
|
||||
# Scheduler 1 raises on every call
|
||||
scheduler1 = MagicMock()
|
||||
scheduler1.execute = AsyncMock(side_effect=RuntimeError("Boom"))
|
||||
|
||||
# Scheduler 2 works normally
|
||||
scheduler2 = MagicMock()
|
||||
scheduler2.execute = AsyncMock()
|
||||
|
||||
async def execute_second(event): # noqa: ARG001
|
||||
processed_second.set()
|
||||
|
||||
scheduler2.execute.side_effect = execute_second
|
||||
|
||||
def get_conf_info(origin):
|
||||
if "event-1" in origin:
|
||||
return {"id": "conf-1", "name": "C1"}
|
||||
return {"id": "conf-2", "name": "C2"}
|
||||
|
||||
mock_config_manager.get_conf_info.side_effect = get_conf_info
|
||||
|
||||
mapping = {"conf-1": scheduler1, "conf-2": scheduler2}
|
||||
event_bus = EventBus(
|
||||
event_queue=event_queue,
|
||||
pipeline_scheduler_mapping=mapping,
|
||||
astrbot_config_mgr=mock_config_manager,
|
||||
)
|
||||
|
||||
# Event 1 should trigger the failing scheduler
|
||||
ev1 = MagicMock()
|
||||
ev1.unified_msg_origin = "event-1"
|
||||
ev1.get_platform_id.return_value = "p1"
|
||||
ev1.get_platform_name.return_value = "P1"
|
||||
ev1.get_sender_name.return_value = None
|
||||
ev1.get_sender_id.return_value = "u1"
|
||||
ev1.get_message_outline.return_value = "m1"
|
||||
|
||||
# Event 2 should trigger the working scheduler
|
||||
ev2 = MagicMock()
|
||||
ev2.unified_msg_origin = "event-2"
|
||||
ev2.get_platform_id.return_value = "p2"
|
||||
ev2.get_platform_name.return_value = "P2"
|
||||
ev2.get_sender_name.return_value = None
|
||||
ev2.get_sender_id.return_value = "u2"
|
||||
ev2.get_message_outline.return_value = "m2"
|
||||
|
||||
await event_queue.put(ev1)
|
||||
await event_queue.put(ev2)
|
||||
|
||||
task = asyncio.create_task(event_bus.dispatch())
|
||||
try:
|
||||
await asyncio.wait_for(processed_second.wait(), timeout=2.0)
|
||||
finally:
|
||||
task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
# Both schedulers should have been called
|
||||
scheduler1.execute.assert_called_once_with(ev1)
|
||||
scheduler2.execute.assert_called_once_with(ev2)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_continues_after_scheduler_exception(
|
||||
self, event_queue, mock_config_manager
|
||||
):
|
||||
"""The dispatch loop must continue processing after a scheduler exception."""
|
||||
scheduler = MagicMock()
|
||||
call_count = 0
|
||||
|
||||
async def execute_alternating(event): # noqa: ARG001
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
raise RuntimeError("First call fails")
|
||||
|
||||
scheduler.execute.side_effect = execute_alternating
|
||||
mock_config_manager.get_conf_info.return_value = {
|
||||
"id": "same-conf",
|
||||
"name": "Same",
|
||||
}
|
||||
mapping = {"same-conf": scheduler}
|
||||
event_bus = EventBus(
|
||||
event_queue=event_queue,
|
||||
pipeline_scheduler_mapping=mapping,
|
||||
astrbot_config_mgr=mock_config_manager,
|
||||
)
|
||||
|
||||
ev = MagicMock()
|
||||
ev.unified_msg_origin = "test:group:1"
|
||||
ev.get_platform_id.return_value = "test"
|
||||
ev.get_platform_name.return_value = "Test"
|
||||
ev.get_sender_name.return_value = None
|
||||
ev.get_sender_id.return_value = "u1"
|
||||
ev.get_message_outline.return_value = "m"
|
||||
|
||||
await event_queue.put(ev)
|
||||
|
||||
task = asyncio.create_task(event_bus.dispatch())
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.sleep(0.2), timeout=1.0)
|
||||
finally:
|
||||
task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
# scheduler should have been called once (event consumed,
|
||||
# exception swallowed by asyncio.create_task)
|
||||
assert call_count >= 1
|
||||
|
||||
|
||||
# ---- Edge-case inputs to dispatch ----
|
||||
|
||||
|
||||
class TestDispatchEdgeInputs:
|
||||
"""Dispatch handles unusual or missing event attributes."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_with_empty_origin(
|
||||
self, event_queue, mock_config_manager, mock_scheduler
|
||||
):
|
||||
"""An event with empty unified_msg_origin is handled (falls back to '')."""
|
||||
processed = asyncio.Event()
|
||||
|
||||
mock_config_manager.get_conf_info.return_value = {
|
||||
"id": "test-conf-id",
|
||||
"name": "Test",
|
||||
}
|
||||
|
||||
async def execute_and_signal(event): # noqa: ARG001
|
||||
processed.set()
|
||||
|
||||
mock_scheduler.execute.side_effect = execute_and_signal
|
||||
|
||||
mapping = {"test-conf-id": mock_scheduler}
|
||||
event_bus = EventBus(
|
||||
event_queue=event_queue,
|
||||
pipeline_scheduler_mapping=mapping,
|
||||
astrbot_config_mgr=mock_config_manager,
|
||||
)
|
||||
|
||||
ev = MagicMock()
|
||||
ev.unified_msg_origin = "" # empty origin
|
||||
ev.get_platform_id.return_value = "test"
|
||||
ev.get_platform_name.return_value = "Test"
|
||||
ev.get_sender_name.return_value = None
|
||||
ev.get_sender_id.return_value = "u1"
|
||||
ev.get_message_outline.return_value = "m"
|
||||
|
||||
await event_queue.put(ev)
|
||||
|
||||
task = asyncio.create_task(event_bus.dispatch())
|
||||
try:
|
||||
await asyncio.wait_for(processed.wait(), timeout=1.0)
|
||||
finally:
|
||||
task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
mock_config_manager.get_conf_info.assert_called_once_with("")
|
||||
mock_scheduler.execute.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_with_empty_config_info(
|
||||
self, event_queue, mock_config_manager, mock_scheduler
|
||||
):
|
||||
"""get_conf_info returning empty dict still works (id falls back to '')."""
|
||||
processed = asyncio.Event()
|
||||
mock_config_manager.get_conf_info.return_value = {}
|
||||
|
||||
async def execute_and_signal(event): # noqa: ARG001
|
||||
processed.set()
|
||||
|
||||
mock_scheduler.execute.side_effect = execute_and_signal
|
||||
|
||||
# The scheduler mapped to '' will be used (empty string from get("id", ""))
|
||||
mapping = {"": mock_scheduler}
|
||||
event_bus = EventBus(
|
||||
event_queue=event_queue,
|
||||
pipeline_scheduler_mapping=mapping,
|
||||
astrbot_config_mgr=mock_config_manager,
|
||||
)
|
||||
|
||||
ev = MagicMock()
|
||||
ev.unified_msg_origin = "some:origin"
|
||||
ev.get_platform_id.return_value = "test"
|
||||
ev.get_platform_name.return_value = "Test"
|
||||
ev.get_sender_name.return_value = None
|
||||
ev.get_sender_id.return_value = "u1"
|
||||
ev.get_message_outline.return_value = "m"
|
||||
|
||||
await event_queue.put(ev)
|
||||
|
||||
task = asyncio.create_task(event_bus.dispatch())
|
||||
try:
|
||||
await asyncio.wait_for(processed.wait(), timeout=1.0)
|
||||
finally:
|
||||
task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
mock_scheduler.execute.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_pipeline_mapping_logs_error(
|
||||
self, event_queue, mock_config_manager
|
||||
):
|
||||
"""An event is dropped with a logged error when no scheduler matches."""
|
||||
error_logged = asyncio.Event()
|
||||
mock_config_manager.get_conf_info.return_value = {
|
||||
"id": "orphan-id",
|
||||
"name": "Orphan",
|
||||
}
|
||||
|
||||
event_bus = EventBus(
|
||||
event_queue=event_queue,
|
||||
pipeline_scheduler_mapping={},
|
||||
astrbot_config_mgr=mock_config_manager,
|
||||
)
|
||||
|
||||
ev = MagicMock()
|
||||
ev.unified_msg_origin = "test:private:1"
|
||||
ev.get_platform_id.return_value = "test"
|
||||
ev.get_platform_name.return_value = "Test"
|
||||
ev.get_sender_name.return_value = "User"
|
||||
ev.get_sender_id.return_value = "u1"
|
||||
ev.get_message_outline.return_value = "m"
|
||||
|
||||
await event_queue.put(ev)
|
||||
|
||||
with patch("astrbot.core.event_bus.logger") as mock_logger:
|
||||
mock_logger.error.side_effect = lambda *a, **kw: error_logged.set()
|
||||
task = asyncio.create_task(event_bus.dispatch())
|
||||
try:
|
||||
await asyncio.wait_for(error_logged.wait(), timeout=1.0)
|
||||
finally:
|
||||
task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
assert "orphan-id" in mock_logger.error.call_args[0][0]
|
||||
|
||||
|
||||
# ---- _print_event edge cases ----
|
||||
|
||||
|
||||
class TestPrintEventEdgeCases:
|
||||
"""_print_event handles unusual sender/platform names."""
|
||||
|
||||
def test_print_event_no_sender_name(self, event_bus_factory):
|
||||
"""When sender_name is None, the log omits the sender name section."""
|
||||
event_bus = event_bus_factory()
|
||||
ev = MagicMock()
|
||||
ev.get_platform_id.return_value = "test"
|
||||
ev.get_platform_name.return_value = "TestPlatform"
|
||||
ev.get_sender_name.return_value = None
|
||||
ev.get_sender_id.return_value = "user-123"
|
||||
ev.get_message_outline.return_value = "Hello"
|
||||
|
||||
with patch("astrbot.core.event_bus.logger") as mock_logger:
|
||||
event_bus._print_event(ev, "MyConfig")
|
||||
|
||||
log_msg = mock_logger.info.call_args[0][0]
|
||||
assert "MyConfig" in log_msg
|
||||
assert "TestPlatform" in log_msg
|
||||
assert "user-123" in log_msg
|
||||
# Without sender name, there should be no '/' before user-123
|
||||
# (the format is [Config] [platform] sender_id: outline)
|
||||
assert "Hello" in log_msg
|
||||
|
||||
def test_print_event_sender_name_is_empty_string(self, event_bus_factory):
|
||||
"""An empty sender_name string is treated similarly to a missing name."""
|
||||
event_bus = event_bus_factory()
|
||||
ev = MagicMock()
|
||||
ev.get_platform_id.return_value = "test"
|
||||
ev.get_platform_name.return_value = "TestPlatform"
|
||||
ev.get_sender_name.return_value = ""
|
||||
ev.get_sender_id.return_value = "user-123"
|
||||
ev.get_message_outline.return_value = "Hello"
|
||||
|
||||
with patch("astrbot.core.event_bus.logger") as mock_logger:
|
||||
event_bus._print_event(ev, "MyConfig")
|
||||
|
||||
mock_logger.info.assert_called_once()
|
||||
|
||||
def test_print_event_with_special_characters(self, event_bus_factory):
|
||||
"""Special characters in event fields do not break logging."""
|
||||
event_bus = event_bus_factory()
|
||||
ev = MagicMock()
|
||||
ev.get_platform_id.return_value = "test@#$"
|
||||
ev.get_platform_name.return_value = "Test[Platform]"
|
||||
ev.get_sender_name.return_value = "User{}|"
|
||||
ev.get_sender_id.return_value = "user:\n456"
|
||||
ev.get_message_outline.return_value = "Hello\nWorld"
|
||||
|
||||
with patch("astrbot.core.event_bus.logger") as mock_logger:
|
||||
event_bus._print_event(ev, "Config[1]")
|
||||
|
||||
mock_logger.info.assert_called_once()
|
||||
|
||||
|
||||
# ---- EventBus construction edge cases ----
|
||||
|
||||
|
||||
class TestEventBusConstructionEdgeCases:
|
||||
"""EventBus handles edge values passed to __init__."""
|
||||
|
||||
def test_empty_pipeline_scheduler_mapping(self, event_queue, mock_config_manager):
|
||||
"""Empty pipeline_mapping is accepted at construction."""
|
||||
bus = EventBus(
|
||||
event_queue=event_queue,
|
||||
pipeline_scheduler_mapping={},
|
||||
astrbot_config_mgr=mock_config_manager,
|
||||
)
|
||||
assert bus.pipeline_scheduler_mapping == {}
|
||||
|
||||
def test_none_in_pipeline_mapping(self, event_queue, mock_config_manager):
|
||||
"""None values in the mapping are accepted (dealt with at dispatch time)."""
|
||||
bus = EventBus(
|
||||
event_queue=event_queue,
|
||||
pipeline_scheduler_mapping={"conf-id": None},
|
||||
astrbot_config_mgr=mock_config_manager,
|
||||
)
|
||||
assert bus.pipeline_scheduler_mapping["conf-id"] is None
|
||||
|
||||
|
||||
# ---- Helper fixtures ----
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def event_bus_factory():
|
||||
"""Factory that creates EventBus with default test dependencies."""
|
||||
|
||||
def _create(event_queue=None, mapping=None, config_mgr=None):
|
||||
import asyncio
|
||||
|
||||
q = event_queue or asyncio.Queue()
|
||||
m = mapping or {"test-id": MagicMock()}
|
||||
cm = config_mgr or MagicMock()
|
||||
return EventBus(
|
||||
event_queue=q,
|
||||
pipeline_scheduler_mapping=m,
|
||||
astrbot_config_mgr=cm,
|
||||
)
|
||||
|
||||
return _create
|
||||
@@ -0,0 +1,267 @@
|
||||
"""Mock-based unit tests for astrbot.core.utils.file_extract.
|
||||
|
||||
All external calls (httpx.AsyncClient and anyio.Path) are mocked so that
|
||||
the tests never touch the network or the filesystem.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from astrbot.core.utils.file_extract import extract_file_moonshotai
|
||||
|
||||
|
||||
def _make_http_status_error(status_code: int = 400) -> httpx.HTTPStatusError:
|
||||
"""Return a realistic HTTPStatusError for use as a side_effect."""
|
||||
request = httpx.Request("POST", "https://api.moonshot.cn/v1/files")
|
||||
response = httpx.Response(status_code, request=request)
|
||||
return httpx.HTTPStatusError(
|
||||
f"HTTP {status_code}",
|
||||
request=request,
|
||||
response=response,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# mocks
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _mock_client() -> tuple[MagicMock, AsyncMock]:
|
||||
"""Create a pre-configured AsyncMock httpx client.
|
||||
|
||||
Returns (mock_client, mock_httpx_cls) where:
|
||||
- mock_httpx_cls.return_value = mock_client
|
||||
- mock_client.__aenter__.return_value = mock_client
|
||||
"""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.__aenter__.return_value = mock_client
|
||||
return mock_client
|
||||
|
||||
|
||||
def _mock_path(file_bytes: bytes = b"dummy file content") -> MagicMock:
|
||||
"""Create a pre-configured MagicMock anyio.Path.
|
||||
|
||||
Returns mock_path where:
|
||||
- mock_path.read_bytes is an AsyncMock returning *file_bytes*
|
||||
"""
|
||||
mock_path = MagicMock(spec=Path) # anyio.Path implements Path-like interface
|
||||
mock_path.name = "test-file.pdf" # used for the upload filename
|
||||
mock_path.read_bytes = AsyncMock(return_value=file_bytes)
|
||||
return mock_path
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestExtractFileMoonshotaiSuccess:
|
||||
|
||||
FAKE_CONTENT = b"%PDF-1.4 fake binary content..."
|
||||
EXTRACTED_TEXT = "Extracted text from Moonshot AI API"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_successful_extraction(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""Happy path: file is uploaded, file-id returned, content fetched."""
|
||||
mock_path = _mock_path(self.FAKE_CONTENT)
|
||||
mock_path_cls.return_value = mock_path
|
||||
|
||||
mock_client = _mock_client()
|
||||
mock_httpx_cls.return_value = mock_client
|
||||
|
||||
upload_resp = MagicMock()
|
||||
upload_resp.json.return_value = {"id": "file-mock-001"}
|
||||
mock_client.post.return_value = upload_resp
|
||||
|
||||
content_resp = MagicMock()
|
||||
content_resp.text = self.EXTRACTED_TEXT
|
||||
mock_client.get.return_value = content_resp
|
||||
|
||||
result = await extract_file_moonshotai("/fake/path/report.pdf", "sk-test-key")
|
||||
|
||||
assert result == self.EXTRACTED_TEXT
|
||||
mock_path_cls.assert_called_once()
|
||||
mock_path.read_bytes.assert_called_once()
|
||||
mock_httpx_cls.assert_called_once()
|
||||
mock_client.post.assert_called_once()
|
||||
mock_client.get.assert_called_once_with("/files/file-mock-001/content")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_client_created_with_authorization_header(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""The Bearer token should match the api_key passed."""
|
||||
mock_path_cls.return_value = _mock_path(b"data")
|
||||
mock_client = _mock_client()
|
||||
mock_httpx_cls.return_value = mock_client
|
||||
mock_client.post.return_value = MagicMock(json=lambda: {"id": "f1"})
|
||||
mock_client.get.return_value = MagicMock(text="ok")
|
||||
|
||||
await extract_file_moonshotai("/f.pdf", "sk-secret-42")
|
||||
|
||||
call_headers = mock_httpx_cls.call_args[1].get("headers", {})
|
||||
assert call_headers.get("Authorization") == "Bearer sk-secret-42"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_passes_file_path_to_anyio_path(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""The file_path argument should be forwarded to anyio.Path()."""
|
||||
mock_path_cls.return_value = _mock_path(b"data")
|
||||
mock_client = _mock_client()
|
||||
mock_httpx_cls.return_value = mock_client
|
||||
mock_client.post.return_value = MagicMock(json=lambda: {"id": "f1"})
|
||||
mock_client.get.return_value = MagicMock(text="ok")
|
||||
|
||||
await extract_file_moonshotai("/my/custom/document.txt", "key")
|
||||
|
||||
args, _ = mock_path_cls.call_args
|
||||
assert str(args[0]) == "/my/custom/document.txt"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Error paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestExtractFileMoonshotaiErrors:
|
||||
|
||||
FAKE_CONTENT = b"fake bytes"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_missing_file_id_raises_value_error(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""Upload response without an 'id' key should raise ValueError."""
|
||||
mock_path_cls.return_value = _mock_path(self.FAKE_CONTENT)
|
||||
mock_client = _mock_client()
|
||||
mock_httpx_cls.return_value = mock_client
|
||||
|
||||
resp = MagicMock()
|
||||
resp.json.return_value = {} # no "id"
|
||||
mock_client.post.return_value = resp
|
||||
|
||||
with pytest.raises(ValueError, match="valid file id"):
|
||||
await extract_file_moonshotai("/f.pdf", "key")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_null_file_id_raises_value_error(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""Upload response with 'id': None should raise ValueError."""
|
||||
mock_path_cls.return_value = _mock_path(self.FAKE_CONTENT)
|
||||
mock_client = _mock_client()
|
||||
mock_httpx_cls.return_value = mock_client
|
||||
|
||||
resp = MagicMock()
|
||||
resp.json.return_value = {"id": None}
|
||||
mock_client.post.return_value = resp
|
||||
|
||||
with pytest.raises(ValueError, match="valid file id"):
|
||||
await extract_file_moonshotai("/f.pdf", "key")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_upload_http_error_propagates(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""Non-2xx status from upload POST should propagate as HTTPStatusError."""
|
||||
mock_path_cls.return_value = _mock_path(self.FAKE_CONTENT)
|
||||
mock_client = _mock_client()
|
||||
mock_httpx_cls.return_value = mock_client
|
||||
|
||||
resp = MagicMock()
|
||||
resp.raise_for_status.side_effect = _make_http_status_error(400)
|
||||
mock_client.post.return_value = resp
|
||||
|
||||
with pytest.raises(httpx.HTTPStatusError):
|
||||
await extract_file_moonshotai("/f.pdf", "key")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_content_fetch_http_error_propagates(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""Non-2xx status from content GET should propagate as HTTPStatusError."""
|
||||
mock_path_cls.return_value = _mock_path(self.FAKE_CONTENT)
|
||||
mock_client = _mock_client()
|
||||
mock_httpx_cls.return_value = mock_client
|
||||
|
||||
upload_resp = MagicMock()
|
||||
upload_resp.json.return_value = {"id": "file-001"}
|
||||
mock_client.post.return_value = upload_resp
|
||||
|
||||
content_resp = MagicMock()
|
||||
content_resp.raise_for_status.side_effect = _make_http_status_error(500)
|
||||
mock_client.get.return_value = content_resp
|
||||
|
||||
with pytest.raises(httpx.HTTPStatusError):
|
||||
await extract_file_moonshotai("/f.pdf", "key")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_connection_error_during_upload_propagates(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""A network-level error on POST should propagate as ConnectError."""
|
||||
mock_path_cls.return_value = _mock_path(self.FAKE_CONTENT)
|
||||
mock_client = _mock_client()
|
||||
mock_httpx_cls.return_value = mock_client
|
||||
|
||||
mock_client.post.side_effect = httpx.ConnectError("connection refused")
|
||||
|
||||
with pytest.raises(httpx.ConnectError):
|
||||
await extract_file_moonshotai("/f.pdf", "key")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("astrbot.core.utils.file_extract.httpx.AsyncClient")
|
||||
@patch("astrbot.core.utils.file_extract.anyio.Path")
|
||||
async def test_file_read_error_prevents_http_call(
|
||||
self,
|
||||
mock_path_cls: MagicMock,
|
||||
mock_httpx_cls: MagicMock,
|
||||
):
|
||||
"""If anyio.Path.read_bytes raises, the function should raise and
|
||||
httpx.AsyncClient should never be constructed."""
|
||||
mock_path = _mock_path(b"ignored")
|
||||
mock_path.read_bytes = AsyncMock(side_effect=OSError(2, "No such file"))
|
||||
mock_path_cls.return_value = mock_path
|
||||
|
||||
with pytest.raises(OSError):
|
||||
await extract_file_moonshotai("/missing/file.txt", "key")
|
||||
|
||||
mock_httpx_cls.assert_not_called()
|
||||
@@ -0,0 +1,314 @@
|
||||
"""
|
||||
Unit tests for FixedSizeChunker.
|
||||
|
||||
Covers construction, chunk method with various inputs (empty text, short text,
|
||||
exact fit, overlap behavior, edge cases with chunk_size/overlap parameters,
|
||||
and kwargs overriding instance defaults).
|
||||
All tests isolate the chunker from any external dependencies.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.knowledge_base.chunking.fixed_size import FixedSizeChunker
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Construction
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFixedSizeChunkerConstruction:
|
||||
"""Test construction of FixedSizeChunker."""
|
||||
|
||||
def test_default_construction(self):
|
||||
"""Test default parameters are set correctly."""
|
||||
chunker = FixedSizeChunker()
|
||||
assert chunker.chunk_size == 512
|
||||
assert chunker.chunk_overlap == 50
|
||||
|
||||
def test_custom_construction(self):
|
||||
"""Test custom parameters are applied."""
|
||||
chunker = FixedSizeChunker(chunk_size=256, chunk_overlap=32)
|
||||
assert chunker.chunk_size == 256
|
||||
assert chunker.chunk_overlap == 32
|
||||
|
||||
def test_zero_overlap_construction(self):
|
||||
"""Test that overlap of 0 is accepted."""
|
||||
chunker = FixedSizeChunker(chunk_size=100, chunk_overlap=0)
|
||||
assert chunker.chunk_size == 100
|
||||
assert chunker.chunk_overlap == 0
|
||||
|
||||
def test_large_overlap_construction(self):
|
||||
"""Test that overlap >= chunk_size is accepted at construction."""
|
||||
chunker = FixedSizeChunker(chunk_size=10, chunk_overlap=10)
|
||||
assert chunker.chunk_size == 10
|
||||
assert chunker.chunk_overlap == 10
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - basic cases
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFixedSizeChunkerChunkBasic:
|
||||
"""Test basic chunk() behavior."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_text_returns_empty_list(self):
|
||||
"""Test that empty text returns an empty list."""
|
||||
chunker = FixedSizeChunker()
|
||||
result = await chunker.chunk("")
|
||||
assert result == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_text_returns_single_chunk(self):
|
||||
"""Test that text shorter than chunk_size returns a single chunk."""
|
||||
chunker = FixedSizeChunker(chunk_size=100)
|
||||
text = "Short text."
|
||||
result = await chunker.chunk(text)
|
||||
assert result == [text]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exact_fit_returns_single_chunk(self):
|
||||
"""Test that text exactly chunk_size returns a single chunk."""
|
||||
chunker = FixedSizeChunker(chunk_size=10)
|
||||
text = "0123456789"
|
||||
result = await chunker.chunk(text)
|
||||
assert result == [text]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_longer_than_chunk_size(self):
|
||||
"""Test that text longer than chunk_size is split into multiple chunks."""
|
||||
chunker = FixedSizeChunker(chunk_size=5, chunk_overlap=0)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
assert result == ["abcde", "fghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_whitespace_only_text(self):
|
||||
"""Test that whitespace-only text is handled without crashing."""
|
||||
chunker = FixedSizeChunker()
|
||||
result = await chunker.chunk(" \n\n ")
|
||||
assert isinstance(result, list)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - overlap behavior
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFixedSizeChunkerOverlap:
|
||||
"""Test overlap behavior in chunk()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_overlap_between_chunks(self):
|
||||
"""Test that chunks overlap correctly."""
|
||||
chunker = FixedSizeChunker(chunk_size=6, chunk_overlap=2)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
# start=0: "abcdef", start=4: "efghij"
|
||||
assert result == ["abcdef", "efghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_overlap_equal_to_chunk_size(self):
|
||||
"""Test that when overlap >= chunk_size, the algorithm prevents an infinite loop."""
|
||||
chunker = FixedSizeChunker(chunk_size=5, chunk_overlap=5)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
# start=0: "abcde", since start >= end (5 >= 5), start becomes end
|
||||
# no infinite loop; next: start=5: "fghij"
|
||||
assert result == ["abcde", "fghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_overlap_greater_than_chunk_size(self):
|
||||
"""Test that when overlap > chunk_size, the algorithm prevents an infinite loop."""
|
||||
chunker = FixedSizeChunker(chunk_size=5, chunk_overlap=10)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
# start=0: "abcde", since start >= end (10 >= 5), start becomes end
|
||||
# start=5: "fghij"
|
||||
assert result == ["abcde", "fghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_overlap(self):
|
||||
"""Test that zero overlap produces non-overlapping chunks."""
|
||||
chunker = FixedSizeChunker(chunk_size=5, chunk_overlap=0)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
assert result == ["abcde", "fghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_partial_last_chunk_with_overlap(self):
|
||||
"""Test the last chunk is included even when shorter than chunk_size."""
|
||||
chunker = FixedSizeChunker(chunk_size=6, chunk_overlap=2)
|
||||
text = "abcdefgh"
|
||||
result = await chunker.chunk(text)
|
||||
# start=0: "abcdef", start=4: "efgh"
|
||||
assert result == ["abcdef", "efgh"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - kwargs override instance defaults
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFixedSizeChunkerKwargs:
|
||||
"""Test that kwargs override instance defaults in chunk()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_size_kwarg_overrides_instance(self):
|
||||
"""Test that chunk_size in kwargs overrides the instance default."""
|
||||
chunker = FixedSizeChunker(chunk_size=500, chunk_overlap=0)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text, chunk_size=5)
|
||||
assert result == ["abcde", "fghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_overlap_kwarg_overrides_instance(self):
|
||||
"""Test that chunk_overlap in kwargs overrides the instance default."""
|
||||
chunker = FixedSizeChunker(chunk_size=6, chunk_overlap=0)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text, chunk_overlap=2)
|
||||
assert result == ["abcdef", "efghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_both_kwargs_override(self):
|
||||
"""Test that both kwargs override instance defaults."""
|
||||
chunker = FixedSizeChunker(chunk_size=100, chunk_overlap=10)
|
||||
text = "abcdefghijklmnopqrstuvwxyz"
|
||||
result = await chunker.chunk(text, chunk_size=10, chunk_overlap=3)
|
||||
# start=0: "abcdefghij", start=7: "hijklmnopq", start=14: "opqrstuvwx",
|
||||
# start=21: "vwxyz"
|
||||
assert len(result) == 4
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - edge cases
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFixedSizeChunkerEdgeCases:
|
||||
"""Test edge cases for chunk()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_character_chunks(self):
|
||||
"""Test with chunk_size=1."""
|
||||
chunker = FixedSizeChunker(chunk_size=1, chunk_overlap=0)
|
||||
text = "abc"
|
||||
result = await chunker.chunk(text)
|
||||
assert result == ["a", "b", "c"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_size_one_with_overlap(self):
|
||||
"""Test with chunk_size=1 and overlap=1 (should not infinite loop)."""
|
||||
chunker = FixedSizeChunker(chunk_size=1, chunk_overlap=1)
|
||||
text = "abc"
|
||||
result = await chunker.chunk(text)
|
||||
# start=0: "a", start>=end (1>=1) -> start=1, "b", ...
|
||||
assert result == ["a", "b", "c"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_large_text(self):
|
||||
"""Test with a large chunk size that covers the entire text."""
|
||||
chunker = FixedSizeChunker(chunk_size=1000, chunk_overlap=50)
|
||||
text = "A" * 900
|
||||
result = await chunker.chunk(text)
|
||||
assert result == [text]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_with_newlines(self):
|
||||
"""Test that newlines are handled as regular characters."""
|
||||
chunker = FixedSizeChunker(chunk_size=10, chunk_overlap=5)
|
||||
text = "line1\nline2\nline3"
|
||||
result = await chunker.chunk(text)
|
||||
assert isinstance(result, list)
|
||||
assert all(isinstance(c, str) for c in result)
|
||||
# All characters should be accounted for
|
||||
assert sum(len(c) for c in result) >= len(text)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unicode_text(self):
|
||||
"""Test that unicode characters are handled correctly."""
|
||||
chunker = FixedSizeChunker(chunk_size=5, chunk_overlap=2)
|
||||
text = "你好世界abc"
|
||||
result = await chunker.chunk(text)
|
||||
# Chinese characters are 1 char each in Python
|
||||
assert isinstance(result, list)
|
||||
assert "".join(result) == text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exact_multiple_of_chunk_size(self):
|
||||
"""Test when text length is an exact multiple of chunk_size (no overlap)."""
|
||||
chunker = FixedSizeChunker(chunk_size=5, chunk_overlap=0)
|
||||
text = "abcde" * 3 # 15 chars
|
||||
result = await chunker.chunk(text)
|
||||
assert result == ["abcde", "fghij", "klmno"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_overlap_larger_than_chunk_minus_one(self):
|
||||
"""Test with overlap = chunk_size - 1 (maximum information overlap)."""
|
||||
chunker = FixedSizeChunker(chunk_size=5, chunk_overlap=4)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
# start=0: "abcde", start=1: "bcdef", ..., start=5: "fghij"
|
||||
assert result == ["abcde", "bcdef", "cdefg", "defgh", "efghi", "fghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ten_thousand_characters(self):
|
||||
"""Test that chunker handles a moderately long string without issue."""
|
||||
chunker = FixedSizeChunker(chunk_size=100, chunk_overlap=10)
|
||||
text = "x" * 10000
|
||||
result = await chunker.chunk(text)
|
||||
# Expected length: ceil((10000 - 10) / (100 - 10)) = ceil(9990/90) = 111
|
||||
assert len(result) == 111
|
||||
assert all(len(c) == 100 or len(c) == 100 for c in result[:-1])
|
||||
assert len(result[-1]) <= 100
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - verification of no content loss
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFixedSizeChunkerContentPreservation:
|
||||
"""Test that chunk() does not lose or reorder content."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_content_loss_without_overlap(self):
|
||||
"""Test no content loss when overlap is 0."""
|
||||
chunker = FixedSizeChunker(chunk_size=10, chunk_overlap=0)
|
||||
text = "0123456789" * 5 # 50 chars
|
||||
result = await chunker.chunk(text)
|
||||
combined = "".join(result)
|
||||
assert combined == text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_content_loss_with_overlap(self):
|
||||
"""Test that with overlap, original content is fully represented."""
|
||||
chunker = FixedSizeChunker(chunk_size=10, chunk_overlap=3)
|
||||
text = "0123456789" * 5
|
||||
result = await chunker.chunk(text)
|
||||
# With overlap, combined will be longer than original
|
||||
combined = "".join(result)
|
||||
assert len(combined) >= len(text)
|
||||
for i, char in enumerate(text):
|
||||
assert char in combined
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunks_are_substrings(self):
|
||||
"""Test that all chunks are valid substrings of the original text."""
|
||||
chunker = FixedSizeChunker(chunk_size=7, chunk_overlap=3)
|
||||
text = "abcdefghijklmnop"
|
||||
result = await chunker.chunk(text)
|
||||
for chunk in result:
|
||||
assert chunk in text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunks_preserve_order(self):
|
||||
"""Test that chunks appear in the same order as the original text."""
|
||||
chunker = FixedSizeChunker(chunk_size=5, chunk_overlap=2)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
# Every chunk should be a substring and order should be consistent
|
||||
positions = [text.index(chunk) for chunk in result]
|
||||
assert positions == sorted(positions)
|
||||
@@ -0,0 +1,926 @@
|
||||
"""
|
||||
Unit tests for KBHelper.
|
||||
|
||||
Covers construction, initialize, get_ep, get_rp, _ensure_vec_db, terminate,
|
||||
delete_vec_db, upload_document, list_documents, get_document, delete_document,
|
||||
delete_chunk, refresh_kb, refresh_document, get_chunks_by_doc_id,
|
||||
get_chunk_count_by_doc_id, _save_media, upload_from_url,
|
||||
and _clean_and_rechunk_content.
|
||||
All tests use mocks to isolate KBHelper from its dependencies.
|
||||
"""
|
||||
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, PropertyMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stub_provider_manager_module():
|
||||
"""Stub provider manager module to avoid circular imports in unit tests."""
|
||||
original_module = sys.modules.get("astrbot.core.provider.manager")
|
||||
stub_module = types.ModuleType("astrbot.core.provider.manager")
|
||||
|
||||
class ProviderManager:
|
||||
...
|
||||
|
||||
setattr(stub_module, "ProviderManager", ProviderManager)
|
||||
sys.modules["astrbot.core.provider.manager"] = stub_module
|
||||
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if original_module is not None:
|
||||
sys.modules["astrbot.core.provider.manager"] = original_module
|
||||
else:
|
||||
sys.modules.pop("astrbot.core.provider.manager", None)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_provider_manager():
|
||||
"""Create a mock ProviderManager."""
|
||||
manager = MagicMock()
|
||||
manager.get_provider_by_id = AsyncMock()
|
||||
manager.acm = MagicMock()
|
||||
manager.acm.default_conf = {}
|
||||
return manager
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_kb_db():
|
||||
"""Create a mock KBSQLiteDatabase."""
|
||||
db = MagicMock()
|
||||
db.get_db = MagicMock()
|
||||
db.list_kbs = AsyncMock(return_value=[])
|
||||
db.get_kb_by_id = AsyncMock()
|
||||
db.list_documents_by_kb = AsyncMock(return_value=[])
|
||||
db.get_document_by_id = AsyncMock()
|
||||
db.delete_document_by_id = AsyncMock()
|
||||
db.update_kb_stats = AsyncMock()
|
||||
return db
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_knowledge_base():
|
||||
"""Create a mock KnowledgeBase instance using lazy import."""
|
||||
from astrbot.core.knowledge_base.models import KnowledgeBase
|
||||
|
||||
kb = KnowledgeBase(
|
||||
kb_name="test_kb",
|
||||
description="Test knowledge base",
|
||||
emoji="test",
|
||||
embedding_provider_id="test-embedding-provider",
|
||||
rerank_provider_id="test-rerank-provider",
|
||||
chunk_size=512,
|
||||
chunk_overlap=50,
|
||||
top_k_dense=50,
|
||||
top_k_sparse=50,
|
||||
top_m_final=5,
|
||||
)
|
||||
return kb
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_embedding_provider():
|
||||
"""Create a mock object that passes isinstance(EmbeddingProvider) checks."""
|
||||
from astrbot.core.provider.provider import EmbeddingProvider
|
||||
|
||||
class FakeEmbeddingProvider(EmbeddingProvider):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
async def get_embedding(self, text: str) -> list[float]:
|
||||
return [0.1, 0.2, 0.3]
|
||||
|
||||
async def get_embeddings(self, text: list[str]) -> list[list[float]]:
|
||||
return [[0.1, 0.2, 0.3] for _ in text]
|
||||
|
||||
def get_dim(self) -> int:
|
||||
return 3
|
||||
|
||||
return FakeEmbeddingProvider()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_rerank_provider():
|
||||
"""Create a mock object that passes isinstance(RerankProvider) checks."""
|
||||
from astrbot.core.provider.provider import RerankProvider
|
||||
|
||||
class FakeRerankProvider(RerankProvider):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
async def rerank_score(self, query: str, documents: list[str]) -> list[float]:
|
||||
return [1.0] * len(documents)
|
||||
|
||||
return FakeRerankProvider()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_chunker():
|
||||
"""Create a mock chunker."""
|
||||
chunker = MagicMock()
|
||||
chunker.chunk = AsyncMock(return_value=["chunk1", "chunk2"])
|
||||
return chunker
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_vec_db():
|
||||
"""Create a mock FaissVecDB."""
|
||||
vec_db = MagicMock()
|
||||
vec_db.initialize = AsyncMock()
|
||||
vec_db.close = AsyncMock()
|
||||
vec_db.insert_batch = AsyncMock()
|
||||
vec_db.delete = AsyncMock()
|
||||
vec_db.count_documents = AsyncMock(return_value=5)
|
||||
vec_db.document_storage = MagicMock()
|
||||
vec_db.document_storage.get_documents = AsyncMock(return_value=[])
|
||||
return vec_db
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def helper_kwargs(
|
||||
mock_kb_db,
|
||||
mock_knowledge_base,
|
||||
mock_provider_manager,
|
||||
mock_chunker,
|
||||
):
|
||||
"""Standard keyword arguments for constructing a KBHelper."""
|
||||
return {
|
||||
"kb_db": mock_kb_db,
|
||||
"kb": mock_knowledge_base,
|
||||
"provider_manager": mock_provider_manager,
|
||||
"kb_root_dir": "/tmp/test_kb_root",
|
||||
"chunker": mock_chunker,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Construction and initialization
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_construction_sets_attributes(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
mock_knowledge_base,
|
||||
mock_provider_manager,
|
||||
mock_chunker,
|
||||
):
|
||||
"""Test that KBHelper.__init__ sets all expected attributes."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper(
|
||||
kb_db=mock_kb_db,
|
||||
kb=mock_knowledge_base,
|
||||
provider_manager=mock_provider_manager,
|
||||
kb_root_dir="/tmp/test_kb_root",
|
||||
chunker=mock_chunker,
|
||||
)
|
||||
|
||||
assert helper.kb_db is mock_kb_db
|
||||
assert helper.kb is mock_knowledge_base
|
||||
assert helper.prov_mgr is mock_provider_manager
|
||||
assert helper.chunker is mock_chunker
|
||||
assert helper.init_error is None
|
||||
assert helper.vec_db is None
|
||||
assert isinstance(helper.kb_dir, Path)
|
||||
assert "medias" in str(helper.kb_medias_dir)
|
||||
assert "files" in str(helper.kb_files_dir)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_calls_ensure_vec_db(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
):
|
||||
"""Test that initialize() delegates to _ensure_vec_db."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
with (
|
||||
patch.object(KBHelper, "_ensure_vec_db", new_callable=AsyncMock) as mock_ensure,
|
||||
):
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
await helper.initialize()
|
||||
|
||||
mock_ensure.assert_awaited_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# get_ep
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_ep_raises_when_no_embedding_provider_id(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
):
|
||||
"""Test that get_ep raises ValueError when kb has no embedding_provider_id."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.kb.embedding_provider_id = None
|
||||
|
||||
with pytest.raises(ValueError, match="未配置 Embedding Provider"):
|
||||
await helper.get_ep()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_ep_raises_when_provider_not_found(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
):
|
||||
"""Test that get_ep raises ValueError when provider is not found by id."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper_kwargs["provider_manager"].get_provider_by_id.return_value = None
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with pytest.raises(ValueError, match="无法找到"):
|
||||
await helper.get_ep()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_ep_raises_when_not_embedding_provider(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that get_ep raises when the returned provider is not EmbeddingProvider."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = MagicMock()
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with pytest.raises(ValueError, match="not an Embedding Provider"):
|
||||
await helper.get_ep()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_ep_returns_embedding_provider(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_embedding_provider,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that get_ep returns the EmbeddingProvider successfully."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = mock_embedding_provider
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
result = await helper.get_ep()
|
||||
assert result is mock_embedding_provider
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# get_rp
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_rp_returns_none_when_no_rerank_provider_id(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
):
|
||||
"""Test that get_rp returns None when kb has no rerank_provider_id."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.kb.rerank_provider_id = None
|
||||
|
||||
result = await helper.get_rp()
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_rp_returns_none_when_provider_not_found(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that get_rp returns None when the provider is not found (logs a warning)."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = None
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
result = await helper.get_rp()
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_rp_raises_when_not_rerank_provider(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that get_rp raises when the returned provider is not RerankProvider."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = MagicMock()
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with pytest.raises(ValueError, match="not a Rerank Provider"):
|
||||
await helper.get_rp()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_rp_returns_rerank_provider(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_rerank_provider,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that get_rp returns the RerankProvider successfully."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = mock_rerank_provider
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
result = await helper.get_rp()
|
||||
assert result is mock_rerank_provider
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# _ensure_vec_db
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_vec_db_creates_and_initializes(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_embedding_provider,
|
||||
mock_provider_manager,
|
||||
mock_vec_db,
|
||||
):
|
||||
"""Test that _ensure_vec_db creates FaissVecDB and initializes it."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = mock_embedding_provider
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with patch(
|
||||
"astrbot.core.knowledge_base.kb_helper.FaissVecDB",
|
||||
return_value=mock_vec_db,
|
||||
) as mock_faiss_cls:
|
||||
result = await helper._ensure_vec_db()
|
||||
|
||||
assert result is mock_vec_db
|
||||
assert helper.vec_db is mock_vec_db
|
||||
assert helper.init_error is None
|
||||
mock_faiss_cls.assert_called_once()
|
||||
mock_vec_db.initialize.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_vec_db_clears_stale_init_error(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_embedding_provider,
|
||||
mock_provider_manager,
|
||||
mock_vec_db,
|
||||
):
|
||||
"""Test that _ensure_vec_db clears a stale init_error on success."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = mock_embedding_provider
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.init_error = "stale error"
|
||||
|
||||
with patch(
|
||||
"astrbot.core.knowledge_base.kb_helper.FaissVecDB",
|
||||
return_value=mock_vec_db,
|
||||
):
|
||||
await helper._ensure_vec_db()
|
||||
assert helper.init_error is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# terminate and delete_vec_db
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminate_closes_vec_db(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_vec_db,
|
||||
):
|
||||
"""Test that terminate() closes the vec_db."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.vec_db = mock_vec_db
|
||||
|
||||
await helper.terminate()
|
||||
mock_vec_db.close.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminate_handles_missing_vec_db(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
):
|
||||
"""Test that terminate() does not crash when vec_db is None."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.vec_db = None
|
||||
|
||||
# Should not raise
|
||||
await helper.terminate()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# upload_document
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_document_with_pre_chunked_text(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
mock_vec_db,
|
||||
mock_embedding_provider,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that upload_document works with pre-chunked text (skips parsing)."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = mock_embedding_provider
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
# Mock the async session for document metadata persistence
|
||||
mock_session = MagicMock()
|
||||
mock_session.add = MagicMock()
|
||||
mock_session.begin = MagicMock()
|
||||
mock_session.begin.return_value.__aenter__ = AsyncMock()
|
||||
mock_session.begin.return_value.__aexit__ = AsyncMock()
|
||||
mock_session.refresh = AsyncMock()
|
||||
mock_db_ctx = MagicMock()
|
||||
mock_db_ctx.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_db_ctx.__aexit__ = AsyncMock()
|
||||
mock_kb_db.get_db.return_value = mock_db_ctx
|
||||
|
||||
with patch.object(helper, "_get_vec_db", return_value=mock_vec_db):
|
||||
doc = await helper.upload_document(
|
||||
file_name="test.txt",
|
||||
file_content=None,
|
||||
file_type="txt",
|
||||
pre_chunked_text=["chunk a", "chunk b"],
|
||||
)
|
||||
|
||||
assert doc.doc_name == "test.txt"
|
||||
assert doc.file_type == "txt"
|
||||
# Metadata should have been added
|
||||
mock_session.add.assert_called()
|
||||
mock_vec_db.insert_batch.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_document_parses_when_no_pre_chunked(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
mock_vec_db,
|
||||
mock_embedding_provider,
|
||||
mock_provider_manager,
|
||||
mock_chunker,
|
||||
):
|
||||
"""Test that upload_document parses file content when no pre-chunked text."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = mock_embedding_provider
|
||||
mock_chunker.chunk = AsyncMock(return_value=["parsed chunk"])
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_session.add = MagicMock()
|
||||
mock_session.begin = MagicMock()
|
||||
mock_session.begin.return_value.__aenter__ = AsyncMock()
|
||||
mock_session.begin.return_value.__aexit__ = AsyncMock()
|
||||
mock_session.refresh = AsyncMock()
|
||||
mock_db_ctx = MagicMock()
|
||||
mock_db_ctx.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_db_ctx.__aexit__ = AsyncMock()
|
||||
mock_kb_db.get_db.return_value = mock_db_ctx
|
||||
|
||||
mock_parser = MagicMock()
|
||||
mock_parse_result = MagicMock()
|
||||
mock_parse_result.text = "some parsed text"
|
||||
mock_parse_result.media = []
|
||||
mock_parser.parse = AsyncMock(return_value=mock_parse_result)
|
||||
|
||||
with (
|
||||
patch.object(helper, "_get_vec_db", return_value=mock_vec_db),
|
||||
patch(
|
||||
"astrbot.core.knowledge_base.kb_helper.select_parser",
|
||||
return_value=mock_parser,
|
||||
),
|
||||
):
|
||||
doc = await helper.upload_document(
|
||||
file_name="test.pdf",
|
||||
file_content=b"%PDF-1.4 fake content",
|
||||
file_type="pdf",
|
||||
)
|
||||
|
||||
assert doc.doc_name == "test.pdf"
|
||||
mock_parser.parse.assert_awaited_once()
|
||||
mock_chunker.chunk.assert_awaited_once()
|
||||
mock_vec_db.insert_batch.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_document_raises_when_no_content_and_no_pre_chunked(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_vec_db,
|
||||
mock_embedding_provider,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that upload_document raises when file_content is None and no pre_chunked."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id.return_value = mock_embedding_provider
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with (
|
||||
patch.object(helper, "_get_vec_db", return_value=mock_vec_db),
|
||||
pytest.raises(ValueError, match="file_content 不能为空"),
|
||||
):
|
||||
await helper.upload_document(
|
||||
file_name="test.pdf",
|
||||
file_content=None,
|
||||
file_type="pdf",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# list_documents / get_document
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_documents_delegates_to_db(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
):
|
||||
"""Test that list_documents delegates to kb_db.list_documents_by_kb."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_kb_db.list_documents_by_kb.return_value = ["doc1", "doc2"]
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
result = await helper.list_documents(offset=0, limit=50)
|
||||
|
||||
mock_kb_db.list_documents_by_kb.assert_awaited_once_with(
|
||||
helper.kb.kb_id,
|
||||
0,
|
||||
50,
|
||||
)
|
||||
assert result == ["doc1", "doc2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_document_delegates_to_db(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
):
|
||||
"""Test that get_document delegates to kb_db.get_document_by_id."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_kb_db.get_document_by_id.return_value = "fake_doc"
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
result = await helper.get_document("doc-123")
|
||||
|
||||
mock_kb_db.get_document_by_id.assert_awaited_once_with("doc-123")
|
||||
assert result == "fake_doc"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# delete_document / delete_chunk
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_document_delegates_to_db_and_vec_db(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
mock_vec_db,
|
||||
):
|
||||
"""Test that delete_document cleans up via kb_db and vec_db."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_kb_db.get_kb_by_id = AsyncMock(return_value=helper_kwargs["kb"])
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.vec_db = mock_vec_db
|
||||
|
||||
await helper.delete_document("doc-123")
|
||||
|
||||
mock_kb_db.delete_document_by_id.assert_awaited_once_with(
|
||||
doc_id="doc-123",
|
||||
vec_db=mock_vec_db,
|
||||
)
|
||||
mock_kb_db.update_kb_stats.assert_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_chunk_delegates_to_vec_db(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
mock_vec_db,
|
||||
):
|
||||
"""Test that delete_chunk removes a chunk from vec_db and refreshes stats."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_kb_db.get_kb_by_id = AsyncMock(return_value=helper_kwargs["kb"])
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.vec_db = mock_vec_db
|
||||
|
||||
await helper.delete_chunk("chunk-1", "doc-123")
|
||||
|
||||
mock_vec_db.delete.assert_awaited_once_with("chunk-1")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# refresh_kb / refresh_document
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_kb_updates_self_kb(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that refresh_kb fetches the latest KB from the database."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
fresh_kb = MagicMock()
|
||||
fresh_kb.kb_id = mock_knowledge_base.kb_id
|
||||
mock_kb_db.get_kb_by_id = AsyncMock(return_value=fresh_kb)
|
||||
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
await helper.refresh_kb()
|
||||
|
||||
mock_kb_db.get_kb_by_id.assert_awaited_once_with(mock_knowledge_base.kb_id)
|
||||
assert helper.kb is fresh_kb
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_document_raises_when_not_found(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_kb_db,
|
||||
):
|
||||
"""Test that refresh_document raises ValueError when doc does not exist."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_kb_db.get_document_by_id = AsyncMock(return_value=None)
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with pytest.raises(ValueError, match="无法找到"):
|
||||
await helper.refresh_document("doc-123")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# get_chunks_by_doc_id
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_chunks_by_doc_id_returns_formatted_chunks(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_vec_db,
|
||||
):
|
||||
"""Test that get_chunks_by_doc_id returns properly formatted chunks."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_vec_db.document_storage.get_documents.return_value = [
|
||||
{
|
||||
"doc_id": "chunk-1",
|
||||
"metadata": json.dumps(
|
||||
{"kb_doc_id": "doc-123", "kb_id": "kb-1", "chunk_index": 0},
|
||||
),
|
||||
"text": "chunk content",
|
||||
},
|
||||
]
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.vec_db = mock_vec_db
|
||||
|
||||
result = await helper.get_chunks_by_doc_id("doc-123")
|
||||
|
||||
assert len(result) == 1
|
||||
chunk = result[0]
|
||||
assert chunk["chunk_id"] == "chunk-1"
|
||||
assert chunk["doc_id"] == "doc-123"
|
||||
assert chunk["content"] == "chunk content"
|
||||
assert chunk["char_count"] == 13
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# get_chunk_count_by_doc_id
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_chunk_count_by_doc_id(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_vec_db,
|
||||
):
|
||||
"""Test that get_chunk_count_by_doc_id returns the count from vec_db."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_vec_db.count_documents.return_value = 42
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
helper.vec_db = mock_vec_db
|
||||
|
||||
count = await helper.get_chunk_count_by_doc_id("doc-123")
|
||||
|
||||
assert count == 42
|
||||
mock_vec_db.count_documents.assert_awaited_once_with(
|
||||
metadata_filter={"kb_doc_id": "doc-123"},
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# _save_media
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_save_media_creates_media_record(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
):
|
||||
"""Test that _save_media saves the media file and returns a KBMedia record."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with patch("astrbot.core.knowledge_base.kb_helper.aiofiles.open") as mock_aiofiles:
|
||||
mock_file = AsyncMock()
|
||||
mock_aiofiles.return_value.__aenter__ = AsyncMock(return_value=mock_file)
|
||||
mock_aiofiles.return_value.__aexit__ = AsyncMock()
|
||||
|
||||
media = await helper._save_media(
|
||||
doc_id="doc-1",
|
||||
media_type="image",
|
||||
file_name="photo.png",
|
||||
content=b"fake-image-bytes",
|
||||
mime_type="image/png",
|
||||
)
|
||||
|
||||
assert media.media_type == "image"
|
||||
assert media.file_name == "photo.png"
|
||||
assert media.mime_type == "image/png"
|
||||
assert media.doc_id == "doc-1"
|
||||
mock_file.write.assert_awaited_once_with(b"fake-image-bytes")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# upload_from_url
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_from_url_raises_when_no_tavily_key(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that upload_from_url raises when Tavily API key is missing."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.acm.default_conf = {"provider_settings": {}}
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with pytest.raises(ValueError, match="Tavily API key"):
|
||||
await helper.upload_from_url("http://example.com")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_from_url_delegates_to_upload_document(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that upload_from_url extracts text and delegates to upload_document."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.acm.default_conf = {
|
||||
"provider_settings": {
|
||||
"websearch_tavily_key": ["fake-key"],
|
||||
},
|
||||
}
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.knowledge_base.kb_helper.extract_text_from_url",
|
||||
new_callable=AsyncMock,
|
||||
return_value="extracted content",
|
||||
),
|
||||
patch.object(
|
||||
helper,
|
||||
"upload_document",
|
||||
new_callable=AsyncMock,
|
||||
return_value="fake_doc",
|
||||
),
|
||||
):
|
||||
result = await helper.upload_from_url("http://example.com")
|
||||
|
||||
assert result == "fake_doc"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# _clean_and_rechunk_content
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clean_and_rechunk_skips_when_not_enabled(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_chunker,
|
||||
):
|
||||
"""Test that _clean_and_rechunk_content uses chunker directly when cleaning disabled."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_chunker.chunk = AsyncMock(return_value=["chunk a", "chunk b"])
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
result = await helper._clean_and_rechunk_content(
|
||||
content="some text",
|
||||
url="http://example.com",
|
||||
enable_cleaning=False,
|
||||
)
|
||||
|
||||
assert result == ["chunk a", "chunk b"]
|
||||
mock_chunker.chunk.assert_awaited_once_with("some text")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clean_and_rechunk_skips_when_no_provider_id(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_chunker,
|
||||
):
|
||||
"""Test that _clean_and_rechunk_content uses chunker when cleaning_provider_id is None."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_chunker.chunk = AsyncMock(return_value=["default chunk"])
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
result = await helper._clean_and_rechunk_content(
|
||||
content="some text",
|
||||
url="http://example.com",
|
||||
enable_cleaning=True,
|
||||
cleaning_provider_id=None,
|
||||
)
|
||||
|
||||
assert result == ["default chunk"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clean_and_rechunk_falls_back_on_error(
|
||||
stub_provider_manager_module,
|
||||
helper_kwargs,
|
||||
mock_provider_manager,
|
||||
mock_chunker,
|
||||
):
|
||||
"""Test that _clean_and_rechunk_content falls back to chunker when LLM call fails."""
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_provider_manager.get_provider_by_id = AsyncMock(
|
||||
side_effect=Exception("provider error"),
|
||||
)
|
||||
mock_chunker.chunk = AsyncMock(return_value=["fallback chunk"])
|
||||
helper = KBHelper(**helper_kwargs)
|
||||
|
||||
result = await helper._clean_and_rechunk_content(
|
||||
content="some text",
|
||||
url="http://example.com",
|
||||
enable_cleaning=True,
|
||||
cleaning_provider_id="llm-1",
|
||||
)
|
||||
|
||||
assert result == ["fallback chunk"]
|
||||
@@ -0,0 +1,657 @@
|
||||
"""
|
||||
Unit tests for KnowledgeBaseManager.
|
||||
|
||||
Covers construction, initialize, create_kb, get_kb, get_kb_by_name,
|
||||
delete_kb, list_kbs, update_kb, retrieve, terminate, and upload_from_url.
|
||||
All tests use mocks to isolate the manager from its dependencies.
|
||||
"""
|
||||
|
||||
import sys
|
||||
import types
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stub_provider_manager_module():
|
||||
"""Stub provider manager module to avoid circular imports in unit tests."""
|
||||
original_module = sys.modules.get("astrbot.core.provider.manager")
|
||||
stub_module = types.ModuleType("astrbot.core.provider.manager")
|
||||
|
||||
class ProviderManager:
|
||||
...
|
||||
|
||||
setattr(stub_module, "ProviderManager", ProviderManager)
|
||||
sys.modules["astrbot.core.provider.manager"] = stub_module
|
||||
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if original_module is not None:
|
||||
sys.modules["astrbot.core.provider.manager"] = original_module
|
||||
else:
|
||||
sys.modules.pop("astrbot.core.provider.manager", None)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_provider_manager():
|
||||
"""Create a mock ProviderManager."""
|
||||
manager = MagicMock()
|
||||
manager.get_provider_by_id = AsyncMock()
|
||||
manager.acm = MagicMock()
|
||||
manager.acm.default_conf = {}
|
||||
return manager
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_kb_db():
|
||||
"""Create a mock KBSQLiteDatabase."""
|
||||
db = MagicMock()
|
||||
db.get_db = MagicMock()
|
||||
db.list_kbs = AsyncMock(return_value=[])
|
||||
db.get_kb_by_id = AsyncMock()
|
||||
return db
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session():
|
||||
"""Create a mock async session with transaction helpers."""
|
||||
session = MagicMock()
|
||||
session.add = MagicMock()
|
||||
session.flush = AsyncMock()
|
||||
session.commit = AsyncMock()
|
||||
session.refresh = AsyncMock()
|
||||
session.begin = MagicMock()
|
||||
session.begin.return_value.__aenter__ = AsyncMock()
|
||||
session.begin.return_value.__aexit__ = AsyncMock()
|
||||
session.delete = AsyncMock()
|
||||
return session
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_db_context(mock_session):
|
||||
"""Create a mock async context manager for get_db()."""
|
||||
ctx = MagicMock()
|
||||
ctx.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
ctx.__aexit__ = AsyncMock()
|
||||
return ctx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_knowledge_base():
|
||||
"""Create a mock KnowledgeBase instance using lazy import."""
|
||||
from astrbot.core.knowledge_base.models import KnowledgeBase
|
||||
|
||||
kb = KnowledgeBase(
|
||||
kb_name="test_kb",
|
||||
description="Test knowledge base",
|
||||
emoji="test",
|
||||
embedding_provider_id="test-embedding-provider",
|
||||
rerank_provider_id=None,
|
||||
chunk_size=512,
|
||||
chunk_overlap=50,
|
||||
top_k_dense=50,
|
||||
top_k_sparse=50,
|
||||
top_m_final=5,
|
||||
)
|
||||
return kb
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Construction & initialization
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_manager_construction(stub_provider_manager_module, mock_provider_manager):
|
||||
"""Test that KnowledgeBaseManager can be constructed with a provider manager."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.provider_manager = mock_provider_manager
|
||||
mgr.kb_insts = {}
|
||||
mgr._session_deleted_callback_registered = False
|
||||
|
||||
assert mgr.provider_manager is mock_provider_manager
|
||||
assert mgr.kb_insts == {}
|
||||
assert mgr._session_deleted_callback_registered is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_manager_initialize_creates_db_and_loads_kbs(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_kb_db,
|
||||
):
|
||||
"""Test that initialize() creates the database and loads existing KBs."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.provider_manager = mock_provider_manager
|
||||
mgr.kb_insts = {}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.knowledge_base.kb_mgr.KBSQLiteDatabase",
|
||||
return_value=mock_kb_db,
|
||||
) as mock_db_cls,
|
||||
patch(
|
||||
"astrbot.core.knowledge_base.kb_mgr.RetrievalManager",
|
||||
) as mock_retrieval_cls,
|
||||
patch(
|
||||
"astrbot.core.knowledge_base.kb_mgr.SparseRetriever",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.knowledge_base.kb_mgr.RankFusion",
|
||||
),
|
||||
):
|
||||
mock_retrieval = MagicMock()
|
||||
mock_retrieval_cls.return_value = mock_retrieval
|
||||
|
||||
await mgr.initialize()
|
||||
|
||||
mock_db_cls.assert_called_once()
|
||||
mock_kb_db.initialize.assert_awaited_once()
|
||||
mock_kb_db.migrate_to_v1.assert_awaited_once()
|
||||
mock_kb_db.list_kbs.assert_awaited_once()
|
||||
assert mgr.retrieval_manager is mock_retrieval
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_handles_import_error_gracefully(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that initialize() catches ImportError without crashing."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.provider_manager = mock_provider_manager
|
||||
mgr.kb_insts = {}
|
||||
|
||||
with patch(
|
||||
"astrbot.core.knowledge_base.kb_mgr.KBSQLiteDatabase",
|
||||
side_effect=ImportError("missing dependency"),
|
||||
):
|
||||
await mgr.initialize()
|
||||
# Should not raise — the error is logged
|
||||
assert not hasattr(mgr, "retrieval_manager") or mgr.retrieval_manager is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# create_kb
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_kb_raises_when_embedding_provider_id_is_none(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_kb_db,
|
||||
):
|
||||
"""Test that create_kb raises ValueError when embedding_provider_id is None."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.provider_manager = mock_provider_manager
|
||||
mgr.kb_db = mock_kb_db
|
||||
mgr.kb_insts = {}
|
||||
|
||||
with pytest.raises(ValueError, match="embedding_provider_id"):
|
||||
await mgr.create_kb(kb_name="my_kb")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_kb_success(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_kb_db,
|
||||
mock_db_context,
|
||||
mock_session,
|
||||
):
|
||||
"""Test that create_kb creates a new KB, persists it, and returns a KBHelper."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_kb_db.get_db.return_value = mock_db_context
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.provider_manager = mock_provider_manager
|
||||
mgr.kb_db = mock_kb_db
|
||||
mgr.kb_insts = {}
|
||||
|
||||
with patch.object(KBHelper, "initialize", new_callable=AsyncMock) as mock_init:
|
||||
mock_init.return_value = None
|
||||
|
||||
result = await mgr.create_kb(
|
||||
kb_name="my_kb",
|
||||
description="desc",
|
||||
emoji="doc",
|
||||
embedding_provider_id="ep-1",
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result.kb.kb_name == "my_kb"
|
||||
assert result.kb.embedding_provider_id == "ep-1"
|
||||
mock_session.add.assert_called_once()
|
||||
mock_session.flush.assert_awaited_once()
|
||||
mock_init.assert_awaited_once()
|
||||
mock_session.commit.assert_awaited_once()
|
||||
assert result.kb.kb_id in mgr.kb_insts
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_kb_duplicate_name_raises(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_kb_db,
|
||||
mock_db_context,
|
||||
mock_session,
|
||||
):
|
||||
"""Test that create_kb raises ValueError on duplicate kb_name."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mock_kb_db.get_db.return_value = mock_db_context
|
||||
# Simulate an IntegrityError-like message
|
||||
mock_session.flush = AsyncMock(
|
||||
side_effect=Exception("UNIQUE constraint failed: kb_name"),
|
||||
)
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.provider_manager = mock_provider_manager
|
||||
mgr.kb_db = mock_kb_db
|
||||
mgr.kb_insts = {}
|
||||
|
||||
with pytest.raises(ValueError, match="已存在"):
|
||||
await mgr.create_kb(
|
||||
kb_name="my_kb",
|
||||
embedding_provider_id="ep-1",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# get_kb / get_kb_by_name
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_kb_returns_none_for_unknown_id(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that get_kb returns None when kb_id is not found."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {}
|
||||
|
||||
result = await mgr.get_kb("nonexistent")
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_kb_returns_helper_for_known_id(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that get_kb returns the correct KBHelper for a known kb_id."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
|
||||
result = await mgr.get_kb(mock_knowledge_base.kb_id)
|
||||
assert result is helper
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_kb_by_name_returns_helper(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that get_kb_by_name returns the correct helper by matching kb_name."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
|
||||
result = await mgr.get_kb_by_name("test_kb")
|
||||
assert result is helper
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_kb_by_name_returns_none_for_missing(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that get_kb_by_name returns None when no match is found."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {}
|
||||
|
||||
result = await mgr.get_kb_by_name("nonexistent")
|
||||
assert result is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# delete_kb
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_kb_returns_false_for_unknown(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that delete_kb returns False when kb_id is not found."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {}
|
||||
|
||||
result = await mgr.delete_kb("nonexistent")
|
||||
assert result is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_kb_removes_helper(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_kb_db,
|
||||
mock_db_context,
|
||||
mock_session,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that delete_kb removes the KBHelper, deletes vec_db, and removes from DB."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
mock_kb_db.get_db.return_value = mock_db_context
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
helper.vec_db = MagicMock()
|
||||
helper.delete_vec_db = AsyncMock()
|
||||
helper.terminate = AsyncMock()
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_db = mock_kb_db
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
|
||||
result = await mgr.delete_kb(mock_knowledge_base.kb_id)
|
||||
|
||||
assert result is True
|
||||
helper.delete_vec_db.assert_awaited_once()
|
||||
mock_session.delete.assert_awaited_once_with(mock_knowledge_base)
|
||||
mock_session.commit.assert_awaited_once()
|
||||
assert mock_knowledge_base.kb_id not in mgr.kb_insts
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# list_kbs
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_kbs_returns_empty_when_no_instances(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that list_kbs returns an empty list when no KBs exist."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {}
|
||||
|
||||
result = await mgr.list_kbs()
|
||||
assert result == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_kbs_returns_all_knowledge_bases(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that list_kbs returns KnowledgeBase objects for all instances."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
|
||||
result = await mgr.list_kbs()
|
||||
assert len(result) == 1
|
||||
assert result[0] is mock_knowledge_base
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# terminate
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminate_closes_all_helpers_and_db(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_kb_db,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that terminate() terminates all helpers and closes the database."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
helper.terminate = AsyncMock()
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_db = mock_kb_db
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
|
||||
mock_kb_db.close = AsyncMock()
|
||||
|
||||
await mgr.terminate()
|
||||
|
||||
helper.terminate.assert_awaited_once()
|
||||
mock_kb_db.close.assert_awaited_once()
|
||||
assert mgr.kb_insts == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminate_handles_helper_failure_gracefully(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_kb_db,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that terminate() continues even if one helper raises."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
helper.terminate = AsyncMock(side_effect=Exception("close failed"))
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_db = mock_kb_db
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
|
||||
mock_kb_db.close = AsyncMock()
|
||||
|
||||
# Should not raise
|
||||
await mgr.terminate()
|
||||
helper.terminate.assert_awaited_once()
|
||||
mock_kb_db.close.assert_awaited_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# retrieve
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_returns_empty_when_no_kb_found(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that retrieve returns an empty dict when no KBs match the given names."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {}
|
||||
|
||||
result = await mgr.retrieve("query", kb_names=["nonexistent"])
|
||||
assert result == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_raises_when_all_kbs_unavailable(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that retrieve raises ValueError when all matching KBs have init_error."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
helper.init_error = "provider unavailable"
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
|
||||
with pytest.raises(ValueError, match="所有请求的知识库均不可用"):
|
||||
await mgr.retrieve("query", kb_names=["test_kb"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve_returns_formatted_results(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that retrieve returns properly formatted results when KBs are available."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
helper.init_error = None
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
mgr.kb_db = MagicMock()
|
||||
|
||||
retrieval_manager = MagicMock()
|
||||
retrieval_manager.retrieve = AsyncMock()
|
||||
|
||||
from collections import namedtuple
|
||||
|
||||
FakeResult = namedtuple(
|
||||
"FakeResult",
|
||||
[
|
||||
"chunk_id",
|
||||
"doc_id",
|
||||
"kb_id",
|
||||
"kb_name",
|
||||
"doc_name",
|
||||
"chunk_index",
|
||||
"content",
|
||||
"score",
|
||||
"metadata",
|
||||
],
|
||||
)
|
||||
fake_result = FakeResult(
|
||||
chunk_id="c1",
|
||||
doc_id="d1",
|
||||
kb_id=mock_knowledge_base.kb_id,
|
||||
kb_name="test_kb",
|
||||
doc_name="doc1.pdf",
|
||||
chunk_index=0,
|
||||
content="some content",
|
||||
score=0.95,
|
||||
metadata={"chunk_index": 0, "char_count": 12},
|
||||
)
|
||||
retrieval_manager.retrieve.return_value = [fake_result]
|
||||
mgr.retrieval_manager = retrieval_manager
|
||||
|
||||
result = await mgr.retrieve("query", kb_names=["test_kb"])
|
||||
|
||||
assert result is not None
|
||||
assert "context_text" in result
|
||||
assert "results" in result
|
||||
assert len(result["results"]) == 1
|
||||
assert result["results"][0]["kb_name"] == "test_kb"
|
||||
assert result["results"][0]["content"] == "some content"
|
||||
retrieval_manager.retrieve.assert_awaited_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# upload_from_url
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_from_url_raises_when_kb_not_found(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
):
|
||||
"""Test that upload_from_url raises ValueError when kb_id is not found."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {}
|
||||
|
||||
with pytest.raises(ValueError, match="not found"):
|
||||
await mgr.upload_from_url("nonexistent", "http://example.com")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upload_from_url_delegates_to_helper(
|
||||
stub_provider_manager_module,
|
||||
mock_provider_manager,
|
||||
mock_knowledge_base,
|
||||
):
|
||||
"""Test that upload_from_url delegates to the correct KBHelper."""
|
||||
from astrbot.core.knowledge_base.kb_mgr import KnowledgeBaseManager
|
||||
from astrbot.core.knowledge_base.kb_helper import KBHelper
|
||||
|
||||
helper = KBHelper.__new__(KBHelper)
|
||||
helper.kb = mock_knowledge_base
|
||||
helper.upload_from_url = AsyncMock(return_value="fake_doc")
|
||||
|
||||
mgr = KnowledgeBaseManager.__new__(KnowledgeBaseManager)
|
||||
mgr.kb_insts = {mock_knowledge_base.kb_id: helper}
|
||||
|
||||
result = await mgr.upload_from_url(
|
||||
mock_knowledge_base.kb_id,
|
||||
"http://example.com",
|
||||
)
|
||||
|
||||
assert result == "fake_doc"
|
||||
helper.upload_from_url.assert_awaited_once_with(
|
||||
url="http://example.com",
|
||||
chunk_size=512,
|
||||
chunk_overlap=50,
|
||||
batch_size=32,
|
||||
tasks_limit=3,
|
||||
max_retries=3,
|
||||
progress_callback=None,
|
||||
)
|
||||
@@ -0,0 +1,283 @@
|
||||
"""Mock-based unit tests for LongTermMemory."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, MagicMock, PropertyMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.builtin_stars.astrbot.long_term_memory import LongTermMemory
|
||||
from astrbot.core.platform.message_type import MessageType
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_acm():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_context():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ltm(mock_acm, mock_context):
|
||||
return LongTermMemory(mock_acm, mock_context)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_event():
|
||||
event = MagicMock()
|
||||
event.unified_msg_origin = "qq:group:123456"
|
||||
event.get_message_type.return_value = MessageType.GROUP_MESSAGE
|
||||
event.is_at_or_wake_command = False
|
||||
event.get_group_id.return_value = "123456"
|
||||
event.get_messages.return_value = []
|
||||
event.message_obj.sender.nickname = "TestUser"
|
||||
return event
|
||||
|
||||
|
||||
class TestLongTermMemoryConstruction:
|
||||
"""Construction and initial state."""
|
||||
|
||||
def test_init_stores_deps(self, ltm, mock_acm, mock_context):
|
||||
assert ltm.acm is mock_acm
|
||||
assert ltm.context is mock_context
|
||||
|
||||
def test_init_session_chats_is_defaultdict(self, ltm):
|
||||
assert ltm.session_chats["any_key"] == []
|
||||
|
||||
@patch.object(LongTermMemory, "cfg", return_value={"max_cnt": 100, "image_caption": False, "enable_active_reply": False})
|
||||
def test_remove_session_returns_count(self, mock_cfg, ltm, mock_event):
|
||||
ltm.session_chats["qq:group:123456"] = ["msg1", "msg2"]
|
||||
cnt = 0
|
||||
import asyncio
|
||||
cnt = asyncio.run(ltm.remove_session(mock_event))
|
||||
assert cnt == 2
|
||||
assert "qq:group:123456" not in ltm.session_chats
|
||||
|
||||
|
||||
class TestCfg:
|
||||
"""Configuration extraction from context."""
|
||||
|
||||
def test_cfg_reads_from_context(self, ltm, mock_context, mock_event):
|
||||
fake_ctx_cfg = {
|
||||
"provider_ltm_settings": {
|
||||
"group_message_max_cnt": "500",
|
||||
"image_caption": True,
|
||||
"image_caption_provider_id": "prov-1",
|
||||
"active_reply": {
|
||||
"enable": True,
|
||||
"method": "possibility_reply",
|
||||
"possibility_reply": 0.5,
|
||||
"whitelist": [],
|
||||
},
|
||||
},
|
||||
"provider_settings": {
|
||||
"image_caption_prompt": "Describe",
|
||||
},
|
||||
}
|
||||
mock_context.get_config.return_value = fake_ctx_cfg
|
||||
result = ltm.cfg(mock_event)
|
||||
assert result["max_cnt"] == 500
|
||||
assert result["image_caption"] is True
|
||||
assert result["enable_active_reply"] is True
|
||||
assert result["image_caption_prompt"] == "Describe"
|
||||
|
||||
def test_cfg_handles_missing_max_cnt(self, ltm, mock_context, mock_event):
|
||||
mock_context.get_config.return_value = {
|
||||
"provider_ltm_settings": {},
|
||||
"provider_settings": {"image_caption_prompt": "x"},
|
||||
}
|
||||
result = ltm.cfg(mock_event)
|
||||
assert result["max_cnt"] == 300 # default fallback
|
||||
assert result["image_caption"] is False
|
||||
|
||||
|
||||
class TestGetImageCaption:
|
||||
"""Image caption fetching."""
|
||||
|
||||
async def test_get_image_caption_uses_using_provider(self, ltm, mock_context):
|
||||
mock_provider = AsyncMock()
|
||||
mock_provider.text_chat = AsyncMock(return_value=MagicMock(completion_text="a cat"))
|
||||
mock_context.get_using_provider.return_value = mock_provider
|
||||
mock_context.get_provider_by_id.return_value = None
|
||||
|
||||
caption = await ltm.get_image_caption("http://img.jpg", "", "Describe")
|
||||
assert caption == "a cat"
|
||||
mock_provider.text_chat.assert_awaited_once()
|
||||
|
||||
async def test_get_image_caption_uses_provider_by_id(self, ltm, mock_context):
|
||||
mock_provider = AsyncMock()
|
||||
mock_provider.text_chat = AsyncMock(return_value=MagicMock(completion_text="a dog"))
|
||||
mock_context.get_provider_by_id.return_value = mock_provider
|
||||
|
||||
caption = await ltm.get_image_caption("http://img.jpg", "custom-provider", "Describe")
|
||||
assert caption == "a dog"
|
||||
|
||||
async def test_get_image_caption_raises_on_missing_provider(self, ltm, mock_context):
|
||||
mock_context.get_provider_by_id.return_value = None
|
||||
with pytest.raises(Exception, match="没有找到 ID 为"):
|
||||
await ltm.get_image_caption("http://img.jpg", "missing-provider", "Describe")
|
||||
|
||||
async def test_get_image_caption_raises_on_non_provider_type(self, ltm, mock_context):
|
||||
mock_context.get_provider_by_id.return_value = "not_a_provider"
|
||||
with pytest.raises(Exception, match="提供商类型错误"):
|
||||
await ltm.get_image_caption("http://img.jpg", "bad-provider", "Describe")
|
||||
|
||||
|
||||
class TestNeedActiveReply:
|
||||
"""Active reply decision logic."""
|
||||
|
||||
async def test_need_active_reply_false_when_disabled(self, ltm, mock_event):
|
||||
with patch.object(LongTermMemory, "cfg", return_value={"enable_active_reply": False}):
|
||||
assert await ltm.need_active_reply(mock_event) is False
|
||||
|
||||
async def test_need_active_reply_false_for_private(self, ltm, mock_event):
|
||||
mock_event.get_message_type.return_value = MessageType.FRIEND_MESSAGE
|
||||
with patch.object(LongTermMemory, "cfg", return_value={"enable_active_reply": True, "ar_whitelist": [], "ar_method": "possibility_reply", "ar_possibility": 1.0}):
|
||||
assert await ltm.need_active_reply(mock_event) is False
|
||||
|
||||
async def test_need_active_reply_false_when_wake_command(self, ltm, mock_event):
|
||||
mock_event.is_at_or_wake_command = True
|
||||
with patch.object(LongTermMemory, "cfg", return_value={"enable_active_reply": True, "ar_whitelist": [], "ar_method": "possibility_reply", "ar_possibility": 1.0}):
|
||||
assert await ltm.need_active_reply(mock_event) is False
|
||||
|
||||
async def test_need_active_reply_whitelist_filters(self, ltm, mock_event):
|
||||
with patch.object(LongTermMemory, "cfg", return_value={
|
||||
"enable_active_reply": True,
|
||||
"ar_whitelist": ["other:group:999"],
|
||||
"ar_method": "possibility_reply",
|
||||
"ar_possibility": 1.0,
|
||||
}):
|
||||
assert await ltm.need_active_reply(mock_event) is False
|
||||
|
||||
async def test_need_active_reply_possibility_triggers(self, ltm, mock_event):
|
||||
with patch.object(LongTermMemory, "cfg", return_value={
|
||||
"enable_active_reply": True,
|
||||
"ar_whitelist": [],
|
||||
"ar_method": "possibility_reply",
|
||||
"ar_possibility": 1.0,
|
||||
}):
|
||||
assert await ltm.need_active_reply(mock_event) is True
|
||||
|
||||
|
||||
class TestHandleMessage:
|
||||
"""Message recording logic."""
|
||||
|
||||
@patch.object(LongTermMemory, "get_image_caption", return_value="sunset")
|
||||
async def test_handle_message_plain_and_image(self, mock_get_caption, ltm, mock_event):
|
||||
from astrbot.api.message_components import Plain, Image
|
||||
|
||||
mock_event.get_messages.return_value = [
|
||||
Plain(text="Hello "),
|
||||
Image(url="http://img.jpg"),
|
||||
Plain(text="world"),
|
||||
]
|
||||
mock_event.get_message_type.return_value = MessageType.GROUP_MESSAGE
|
||||
with patch.object(LongTermMemory, "cfg", return_value={
|
||||
"max_cnt": 300,
|
||||
"image_caption": True,
|
||||
"image_caption_provider_id": "prov-1",
|
||||
"image_caption_prompt": "Describe",
|
||||
"enable_active_reply": False,
|
||||
}):
|
||||
await ltm.handle_message(mock_event)
|
||||
assert len(ltm.session_chats["qq:group:123456"]) == 1
|
||||
msg = ltm.session_chats["qq:group:123456"][0]
|
||||
assert "Hello" in msg
|
||||
assert "sunset" in msg
|
||||
assert "world" in msg
|
||||
|
||||
async def test_handle_message_plain_only(self, ltm, mock_event):
|
||||
from astrbot.api.message_components import Plain
|
||||
|
||||
mock_event.get_messages.return_value = [
|
||||
Plain(text="How are you?"),
|
||||
]
|
||||
mock_event.get_message_type.return_value = MessageType.GROUP_MESSAGE
|
||||
with patch.object(LongTermMemory, "cfg", return_value={
|
||||
"max_cnt": 300,
|
||||
"image_caption": False,
|
||||
"enable_active_reply": False,
|
||||
}):
|
||||
await ltm.handle_message(mock_event)
|
||||
msg = ltm.session_chats["qq:group:123456"][0]
|
||||
assert "How are you?" in msg
|
||||
|
||||
async def test_handle_message_ignores_private(self, ltm, mock_event):
|
||||
mock_event.get_message_type.return_value = MessageType.FRIEND_MESSAGE
|
||||
await ltm.handle_message(mock_event)
|
||||
assert "qq:group:123456" not in ltm.session_chats
|
||||
|
||||
async def test_handle_message_trims_excess(self, ltm, mock_event):
|
||||
from astrbot.api.message_components import Plain
|
||||
|
||||
mock_event.get_messages.return_value = [Plain(text="x")]
|
||||
with patch.object(LongTermMemory, "cfg", return_value={
|
||||
"max_cnt": 2,
|
||||
"image_caption": False,
|
||||
"enable_active_reply": False,
|
||||
}):
|
||||
ltm.session_chats["qq:group:123456"] = ["old1", "old2"]
|
||||
await ltm.handle_message(mock_event)
|
||||
msgs = ltm.session_chats["qq:group:123456"]
|
||||
assert len(msgs) == 2
|
||||
|
||||
|
||||
class TestOnReqLLM:
|
||||
"""LLM request modification."""
|
||||
|
||||
async def test_on_req_llm_skips_when_no_history(self, ltm, mock_event):
|
||||
req = MagicMock()
|
||||
await ltm.on_req_llm(mock_event, req)
|
||||
req.assert_not_called()
|
||||
|
||||
async def test_on_req_llm_adds_to_system_prompt(self, ltm, mock_event):
|
||||
ltm.session_chats["qq:group:123456"] = ["msg1", "msg2"]
|
||||
req = MagicMock()
|
||||
req.system_prompt = ""
|
||||
with patch.object(LongTermMemory, "cfg", return_value={"enable_active_reply": False}):
|
||||
await ltm.on_req_llm(mock_event, req)
|
||||
assert "msg1" in req.system_prompt
|
||||
assert "msg2" in req.system_prompt
|
||||
|
||||
async def test_on_req_llm_active_reply_overrides_prompt(self, ltm, mock_event):
|
||||
ltm.session_chats["qq:group:123456"] = ["msg1"]
|
||||
req = MagicMock()
|
||||
req.prompt = "user query"
|
||||
req.contexts = ["old_ctx"]
|
||||
with patch.object(LongTermMemory, "cfg", return_value={"enable_active_reply": True}):
|
||||
await ltm.on_req_llm(mock_event, req)
|
||||
assert req.contexts == []
|
||||
assert "user query" in req.prompt
|
||||
|
||||
|
||||
class TestAfterReqLLM:
|
||||
"""Post-LLM response recording."""
|
||||
|
||||
async def test_after_req_llm_records_response(self, ltm, mock_event):
|
||||
ltm.session_chats["qq:group:123456"] = []
|
||||
resp = MagicMock()
|
||||
resp.completion_text = "AI reply"
|
||||
with patch.object(LongTermMemory, "cfg", return_value={"max_cnt": 300}):
|
||||
await ltm.after_req_llm(mock_event, resp)
|
||||
stored = ltm.session_chats["qq:group:123456"]
|
||||
assert len(stored) == 1
|
||||
assert "AI reply" in stored[0]
|
||||
|
||||
async def test_after_req_llm_skips_empty_response(self, ltm, mock_event):
|
||||
ltm.session_chats["qq:group:123456"] = ["existing"]
|
||||
resp = MagicMock()
|
||||
resp.completion_text = None
|
||||
await ltm.after_req_llm(mock_event, resp)
|
||||
assert len(ltm.session_chats["qq:group:123456"]) == 1
|
||||
|
||||
async def test_after_req_llm_trims(self, ltm, mock_event):
|
||||
ltm.session_chats["qq:group:123456"] = ["m1", "m2", "m3"]
|
||||
resp = MagicMock()
|
||||
resp.completion_text = "new"
|
||||
with patch.object(LongTermMemory, "cfg", return_value={"max_cnt": 3}):
|
||||
await ltm.after_req_llm(mock_event, resp)
|
||||
assert len(ltm.session_chats["qq:group:123456"]) == 3
|
||||
@@ -0,0 +1,335 @@
|
||||
"""Tests for astrbot.core.platform.manager — PlatformManager."""
|
||||
|
||||
import asyncio
|
||||
from asyncio import Queue
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.platform.manager import PlatformManager
|
||||
from astrbot.core.platform.platform import Platform, PlatformError, PlatformStatus
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# A minimal concrete Platform we can wire into the manager without real I/O.
|
||||
# ---------------------------------------------------------------------------
|
||||
class DummyPlatform(Platform):
|
||||
"""Lightweight Platform used inside manager tests."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: dict,
|
||||
settings: dict,
|
||||
event_queue: Queue,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
super().__init__(config, event_queue)
|
||||
self.settings = settings
|
||||
self._terminated = False
|
||||
|
||||
async def run(self) -> None:
|
||||
pass
|
||||
|
||||
def meta(self):
|
||||
from astrbot.core.platform.platform_metadata import PlatformMetadata
|
||||
return PlatformMetadata(
|
||||
name=self.config.get("type", "dummy"),
|
||||
description="dummy",
|
||||
id=self.config.get("id", "dummy_id"),
|
||||
)
|
||||
|
||||
async def terminate(self) -> None:
|
||||
self._terminated = True
|
||||
await super().terminate()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.fixture
|
||||
def mock_config() -> MagicMock:
|
||||
"""A dict-like MagicMock that behaves like AstrBotConfig for key lookups."""
|
||||
cfg = MagicMock()
|
||||
cfg.__getitem__.side_effect = lambda k: {
|
||||
"platform": [],
|
||||
"platform_settings": {"unique_session": True},
|
||||
}[k]
|
||||
return cfg
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def event_queue() -> Queue:
|
||||
return Queue()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def manager(mock_config: MagicMock, event_queue: Queue) -> PlatformManager:
|
||||
return PlatformManager(mock_config, event_queue)
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Construction
|
||||
# ===================================================================
|
||||
|
||||
class TestConstruction:
|
||||
"""PlatformManager.__init__ stores constructor arguments and initialises
|
||||
empty collections."""
|
||||
|
||||
def test_stores_config(self, mock_config: MagicMock, event_queue: Queue):
|
||||
m = PlatformManager(mock_config, event_queue)
|
||||
assert m.astrbot_config is mock_config
|
||||
assert m.event_queue is event_queue
|
||||
|
||||
def test_platforms_config_and_settings_are_extracted(
|
||||
self, mock_config: MagicMock, event_queue: Queue
|
||||
):
|
||||
m = PlatformManager(mock_config, event_queue)
|
||||
assert m.platforms_config == []
|
||||
assert m.settings == {"unique_session": True}
|
||||
|
||||
def test_initial_collections_are_empty(
|
||||
self, mock_config: MagicMock, event_queue: Queue
|
||||
):
|
||||
m = PlatformManager(mock_config, event_queue)
|
||||
assert m.platform_insts == []
|
||||
assert m._inst_map == {}
|
||||
assert m._platform_tasks == {}
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# _is_valid_platform_id
|
||||
# ===================================================================
|
||||
|
||||
class TestIsValidPlatformId:
|
||||
"""_is_valid_platform_id rejects None, empty, or ids containing ':'/'!'."""
|
||||
|
||||
def test_valid_id(self, manager: PlatformManager):
|
||||
assert manager._is_valid_platform_id("my_platform") is True
|
||||
assert manager._is_valid_platform_id("platform123") is True
|
||||
|
||||
def test_none_is_invalid(self, manager: PlatformManager):
|
||||
assert manager._is_valid_platform_id(None) is False
|
||||
|
||||
def test_empty_string_is_invalid(self, manager: PlatformManager):
|
||||
assert manager._is_valid_platform_id("") is False
|
||||
|
||||
def test_colon_is_invalid(self, manager: PlatformManager):
|
||||
assert manager._is_valid_platform_id("plat:form") is False
|
||||
|
||||
def test_exclamation_is_invalid(self, manager: PlatformManager):
|
||||
assert manager._is_valid_platform_id("plat!form") is False
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# _sanitize_platform_id
|
||||
# ===================================================================
|
||||
|
||||
class TestSanitizePlatformId:
|
||||
"""_sanitize_platform_id replaces ':'/'!' with '_'."""
|
||||
|
||||
def test_clean_id_unchanged(self, manager: PlatformManager):
|
||||
result, changed = manager._sanitize_platform_id("my_platform")
|
||||
assert result == "my_platform"
|
||||
assert changed is False
|
||||
|
||||
def test_colon_replaced(self, manager: PlatformManager):
|
||||
result, changed = manager._sanitize_platform_id("my:platform")
|
||||
assert result == "my_platform"
|
||||
assert changed is True
|
||||
|
||||
def test_exclamation_replaced(self, manager: PlatformManager):
|
||||
result, changed = manager._sanitize_platform_id("my!platform")
|
||||
assert result == "my_platform"
|
||||
assert changed is True
|
||||
|
||||
def test_both_replaced(self, manager: PlatformManager):
|
||||
result, changed = manager._sanitize_platform_id("a:b!c")
|
||||
assert result == "a_b_c"
|
||||
assert changed is True
|
||||
|
||||
def test_none_returns_none_no_change(self, manager: PlatformManager):
|
||||
result, changed = manager._sanitize_platform_id(None)
|
||||
assert result is None
|
||||
assert changed is False
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# get_insts
|
||||
# ===================================================================
|
||||
|
||||
class TestGetInsts:
|
||||
"""get_insts returns the internal platform_insts list."""
|
||||
|
||||
def test_returns_same_list_object(self, manager: PlatformManager):
|
||||
assert manager.get_insts() is manager.platform_insts
|
||||
|
||||
def test_empty_by_default(self, manager: PlatformManager):
|
||||
assert manager.get_insts() == []
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# get_all_stats
|
||||
# ===================================================================
|
||||
|
||||
class TestGetAllStats:
|
||||
"""get_all_stats aggregates stats from all registered platforms."""
|
||||
|
||||
def test_empty_when_no_platforms(self, manager: PlatformManager):
|
||||
stats = manager.get_all_stats()
|
||||
assert stats["platforms"] == []
|
||||
assert stats["summary"]["total"] == 0
|
||||
assert stats["summary"]["running"] == 0
|
||||
assert stats["summary"]["error"] == 0
|
||||
assert stats["summary"]["total_errors"] == 0
|
||||
|
||||
def test_aggregates_mixed_statuses(self, manager: PlatformManager):
|
||||
run_mock = MagicMock(spec=Platform)
|
||||
run_mock.get_stats.return_value = {
|
||||
"id": "p1",
|
||||
"type": "mock",
|
||||
"display_name": "Mock1",
|
||||
"status": PlatformStatus.RUNNING.value,
|
||||
"started_at": None,
|
||||
"error_count": 0,
|
||||
"last_error": None,
|
||||
"unified_webhook": False,
|
||||
"meta": {},
|
||||
}
|
||||
err_mock = MagicMock(spec=Platform)
|
||||
err_mock.get_stats.return_value = {
|
||||
"id": "p2",
|
||||
"type": "mock",
|
||||
"display_name": "Mock2",
|
||||
"status": PlatformStatus.ERROR.value,
|
||||
"started_at": None,
|
||||
"error_count": 2,
|
||||
"last_error": {"message": "fail"},
|
||||
"unified_webhook": False,
|
||||
"meta": {},
|
||||
}
|
||||
manager.platform_insts = [run_mock, err_mock]
|
||||
|
||||
stats = manager.get_all_stats()
|
||||
assert stats["summary"]["total"] == 2
|
||||
assert stats["summary"]["running"] == 1
|
||||
assert stats["summary"]["error"] == 1
|
||||
assert stats["summary"]["total_errors"] == 2
|
||||
|
||||
def test_recovers_from_broken_platform(self, manager: PlatformManager):
|
||||
bad = MagicMock(spec=Platform)
|
||||
bad.get_stats.side_effect = Exception("oops")
|
||||
bad.config = {"id": "broken"}
|
||||
manager.platform_insts = [bad]
|
||||
|
||||
stats = manager.get_all_stats()
|
||||
assert stats["summary"]["total"] == 1
|
||||
platform_info = stats["platforms"][0]
|
||||
assert platform_info["id"] == "broken"
|
||||
assert platform_info["status"] == "unknown"
|
||||
assert platform_info["type"] == "unknown"
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# terminate_platform
|
||||
# ===================================================================
|
||||
|
||||
class TestTerminatePlatform:
|
||||
"""terminate_platform removes the platform from maps and calls terminate."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_removes_and_terminates(self, manager: PlatformManager):
|
||||
p = DummyPlatform(
|
||||
{"id": "test_id", "type": "dummy"}, {}, Queue()
|
||||
)
|
||||
manager._inst_map["test_id"] = {
|
||||
"inst": p,
|
||||
"client_id": p.client_self_id,
|
||||
}
|
||||
manager.platform_insts.append(p)
|
||||
|
||||
# Stop the underlying task machinery from actually creating tasks.
|
||||
with patch.object(manager, "_stop_platform_task", new_callable=AsyncMock):
|
||||
await manager.terminate_platform("test_id")
|
||||
|
||||
assert p._terminated is True
|
||||
assert "test_id" not in manager._inst_map
|
||||
assert p not in manager.platform_insts
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# terminate (all)
|
||||
# ===================================================================
|
||||
|
||||
class TestTerminateAll:
|
||||
"""terminate() cleans up all registered platforms."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminates_all_platforms(self, manager: PlatformManager):
|
||||
p1 = DummyPlatform({"id": "p1"}, {}, Queue())
|
||||
p2 = DummyPlatform({"id": "p2"}, {}, Queue())
|
||||
manager._inst_map["p1"] = {"inst": p1, "client_id": p1.client_self_id}
|
||||
manager._inst_map["p2"] = {"inst": p2, "client_id": p2.client_self_id}
|
||||
manager.platform_insts = [p1, p2]
|
||||
|
||||
with patch.object(manager, "_stop_platform_task", new_callable=AsyncMock):
|
||||
await manager.terminate()
|
||||
|
||||
assert p1._terminated is True
|
||||
assert p2._terminated is True
|
||||
assert manager.platform_insts == []
|
||||
assert manager._inst_map == {}
|
||||
assert manager._platform_tasks == {}
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# load_platform
|
||||
# ===================================================================
|
||||
|
||||
class TestLoadPlatform:
|
||||
"""load_platform handles disabled flag and invalid IDs."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_disabled_platform(self, manager: PlatformManager):
|
||||
"""When enable is False, load_platform should return immediately."""
|
||||
with patch("astrbot.core.platform.manager.logger") as mock_logger:
|
||||
await manager.load_platform({"enable": False})
|
||||
# info should not have been called with the loading message
|
||||
for call in mock_logger.info.call_args_list:
|
||||
assert "Loading" not in str(call)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sanitizes_invalid_platform_id(
|
||||
self, manager: PlatformManager
|
||||
):
|
||||
"""A platform ID containing ':' should be sanitized automatically."""
|
||||
platform_cfg = {
|
||||
"enable": True,
|
||||
"type": "nonexistent_type",
|
||||
"id": "bad:id",
|
||||
}
|
||||
with (
|
||||
patch.object(manager.astrbot_config, "save_config"),
|
||||
patch("astrbot.core.platform.manager.logger"),
|
||||
):
|
||||
await manager.load_platform(platform_cfg)
|
||||
|
||||
assert platform_cfg["id"] == "bad_id"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logs_error_when_type_not_in_map(
|
||||
self, manager: PlatformManager
|
||||
):
|
||||
"""If the platform type is not registered, load_platform logs an error."""
|
||||
platform_cfg = {
|
||||
"enable": True,
|
||||
"type": "no_such_adapter",
|
||||
"id": "test_id",
|
||||
}
|
||||
with patch("astrbot.core.platform.manager.logger") as mock_logger:
|
||||
await manager.load_platform(platform_cfg)
|
||||
# Should have logged an error about adapter not found
|
||||
error_messages = [
|
||||
str(c) for c in mock_logger.error.call_args_list
|
||||
]
|
||||
assert any("not found" in msg for msg in error_messages)
|
||||
@@ -0,0 +1,169 @@
|
||||
"""Tests for astrbot.core.platform.message_session — MessageSession."""
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.platform.message_session import MessageSession, MessageSesion
|
||||
from astrbot.core.platform.message_type import MessageType
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Construction
|
||||
# ===================================================================
|
||||
|
||||
class TestConstruction:
|
||||
"""MessageSession construction and post_init."""
|
||||
|
||||
def test_basic_construction(self):
|
||||
session = MessageSession(
|
||||
platform_name="test_platform",
|
||||
message_type=MessageType.FRIEND_MESSAGE,
|
||||
session_id="session_123",
|
||||
)
|
||||
assert session.platform_name == "test_platform"
|
||||
assert session.message_type == MessageType.FRIEND_MESSAGE
|
||||
assert session.session_id == "session_123"
|
||||
|
||||
def test_post_init_sets_platform_id_from_platform_name(self):
|
||||
"""platform_id should be auto-set to platform_name in __post_init__."""
|
||||
session = MessageSession(
|
||||
platform_name="my_adapter",
|
||||
message_type=MessageType.GROUP_MESSAGE,
|
||||
session_id="g123",
|
||||
)
|
||||
assert session.platform_id == "my_adapter"
|
||||
|
||||
def test_platform_id_equals_platform_name(self):
|
||||
"""Confirm platform_id and platform_name hold the same value."""
|
||||
session = MessageSession(
|
||||
platform_name="pname",
|
||||
message_type=MessageType.OTHER_MESSAGE,
|
||||
session_id="sid",
|
||||
)
|
||||
assert session.platform_id == session.platform_name
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# __str__
|
||||
# ===================================================================
|
||||
|
||||
class TestStr:
|
||||
"""MessageSession.__str__ produces the unified-msg-origin format."""
|
||||
|
||||
def test_friend_message_format(self):
|
||||
session = MessageSession(
|
||||
platform_name="discord",
|
||||
message_type=MessageType.FRIEND_MESSAGE,
|
||||
session_id="user_42",
|
||||
)
|
||||
assert str(session) == "discord:FriendMessage:user_42"
|
||||
|
||||
def test_group_message_format(self):
|
||||
session = MessageSession(
|
||||
platform_name="telegram",
|
||||
message_type=MessageType.GROUP_MESSAGE,
|
||||
session_id="group_99",
|
||||
)
|
||||
assert str(session) == "telegram:GroupMessage:group_99"
|
||||
|
||||
def test_other_message_format(self):
|
||||
session = MessageSession(
|
||||
platform_name="system",
|
||||
message_type=MessageType.OTHER_MESSAGE,
|
||||
session_id="sys_1",
|
||||
)
|
||||
assert str(session) == "system:OtherMessage:sys_1"
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# from_str
|
||||
# ===================================================================
|
||||
|
||||
class TestFromStr:
|
||||
"""MessageSession.from_str parses the unified-msg-origin string."""
|
||||
|
||||
def test_parses_friend_message(self):
|
||||
session = MessageSession.from_str("qq:FriendMessage:user_007")
|
||||
assert session.platform_name == "qq"
|
||||
assert session.platform_id == "qq"
|
||||
assert session.message_type == MessageType.FRIEND_MESSAGE
|
||||
assert session.session_id == "user_007"
|
||||
|
||||
def test_parses_group_message(self):
|
||||
session = MessageSession.from_str(
|
||||
"slack:GroupMessage:channel_C01"
|
||||
)
|
||||
assert session.platform_name == "slack"
|
||||
assert session.message_type == MessageType.GROUP_MESSAGE
|
||||
assert session.session_id == "channel_C01"
|
||||
|
||||
@staticmethod
|
||||
def test_roundtrip():
|
||||
"""str -> from_str -> str should be lossless."""
|
||||
original = "wechat:FriendMessage:wx_abc"
|
||||
session = MessageSession.from_str(original)
|
||||
assert str(session) == original
|
||||
|
||||
@staticmethod
|
||||
def test_session_id_contains_colons():
|
||||
"""If the session_id itself contains colons, only the first two split
|
||||
boundaries are consumed; the rest is part of session_id."""
|
||||
raw = "platform:GroupMessage:user:name:extra"
|
||||
session = MessageSession.from_str(raw)
|
||||
assert session.platform_name == "platform"
|
||||
assert session.message_type == MessageType.GROUP_MESSAGE
|
||||
assert session.session_id == "user:name:extra"
|
||||
assert str(session) == raw
|
||||
|
||||
@staticmethod
|
||||
def test_preserves_explicit_platform_id_after_roundtrip():
|
||||
"""from_str sets platform_name from the parsed platform_id segment,
|
||||
so platform_id == platform_name after a roundtrip."""
|
||||
session = MessageSession.from_str("my_id:FriendMessage:sid")
|
||||
assert session.platform_name == "my_id"
|
||||
assert session.platform_id == "my_id"
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Back-compat alias
|
||||
# ===================================================================
|
||||
|
||||
class TestAlias:
|
||||
"""MessageSesion (note the typo) should be an alias for MessageSession."""
|
||||
|
||||
def test_alias_is_same_class(self):
|
||||
assert MessageSesion is MessageSession
|
||||
|
||||
def test_alias_can_be_instantiated(self):
|
||||
session = MessageSesion(
|
||||
platform_name="alias_test",
|
||||
message_type=MessageType.GROUP_MESSAGE,
|
||||
session_id="alias_sid",
|
||||
)
|
||||
assert isinstance(session, MessageSession)
|
||||
assert session.platform_name == "alias_test"
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Dataclass equality
|
||||
# ===================================================================
|
||||
|
||||
class TestEquality:
|
||||
"""MessageSession is a dataclass so it inherits __eq__."""
|
||||
|
||||
def test_equal_sessions(self):
|
||||
a = MessageSession(
|
||||
platform_name="p", message_type=MessageType.FRIEND_MESSAGE, session_id="s"
|
||||
)
|
||||
b = MessageSession(
|
||||
platform_name="p", message_type=MessageType.FRIEND_MESSAGE, session_id="s"
|
||||
)
|
||||
assert a == b
|
||||
|
||||
def test_inequality_different_session_id(self):
|
||||
a = MessageSession(
|
||||
platform_name="p", message_type=MessageType.FRIEND_MESSAGE, session_id="s1"
|
||||
)
|
||||
b = MessageSession(
|
||||
platform_name="p", message_type=MessageType.FRIEND_MESSAGE, session_id="s2"
|
||||
)
|
||||
assert a != b
|
||||
@@ -1,8 +1,27 @@
|
||||
import ssl
|
||||
"""Unit tests for astrbot.core.utils.network_utils.
|
||||
|
||||
Expands upon the existing 2 monkeypatch-based tests with comprehensive
|
||||
coverage of is_connection_error, log_connection_failure, and
|
||||
create_proxy_client edge cases.
|
||||
"""
|
||||
|
||||
import os
|
||||
import ssl
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from astrbot.core.utils import network_utils
|
||||
from astrbot.core.utils import network_utils as network_utils_module
|
||||
from astrbot.core.utils.network_utils import (
|
||||
create_proxy_client,
|
||||
is_connection_error,
|
||||
log_connection_failure,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Existing tests (preserved verbatim)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_create_proxy_client_reuses_shared_ssl_context(
|
||||
@@ -15,12 +34,12 @@ def test_create_proxy_client_reuses_shared_ssl_context(
|
||||
def __init__(self, **kwargs):
|
||||
captured_calls.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(network_utils.httpx, "AsyncClient", _FakeAsyncClient)
|
||||
monkeypatch.setattr(network_utils_module.httpx, "AsyncClient", _FakeAsyncClient)
|
||||
|
||||
network_utils.create_proxy_client("OpenAI")
|
||||
network_utils.create_proxy_client("OpenAI", proxy="http://127.0.0.1:7890")
|
||||
network_utils.create_proxy_client("OpenAI", headers=headers)
|
||||
network_utils.create_proxy_client("OpenAI", proxy="")
|
||||
network_utils_module.create_proxy_client("OpenAI")
|
||||
network_utils_module.create_proxy_client("OpenAI", proxy="http://127.0.0.1:7890")
|
||||
network_utils_module.create_proxy_client("OpenAI", headers=headers)
|
||||
network_utils_module.create_proxy_client("OpenAI", proxy="")
|
||||
|
||||
assert len(captured_calls) == 4
|
||||
assert "proxy" not in captured_calls[0]
|
||||
@@ -43,9 +62,172 @@ def test_create_proxy_client_allows_verify_override(
|
||||
def __init__(self, **kwargs):
|
||||
captured_calls.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(network_utils.httpx, "AsyncClient", _FakeAsyncClient)
|
||||
monkeypatch.setattr(network_utils_module.httpx, "AsyncClient", _FakeAsyncClient)
|
||||
|
||||
network_utils.create_proxy_client("OpenAI", verify=custom_verify)
|
||||
network_utils_module.create_proxy_client("OpenAI", verify=custom_verify)
|
||||
|
||||
assert len(captured_calls) == 1
|
||||
assert captured_calls[0]["verify"] is custom_verify
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# is_connection_error
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestIsConnectionError:
|
||||
def test_httpx_connect_error(self):
|
||||
assert is_connection_error(httpx.ConnectError("refused")) is True
|
||||
|
||||
def test_httpx_timeout_errors(self):
|
||||
for exc_cls in (
|
||||
httpx.ConnectTimeout,
|
||||
httpx.ReadTimeout,
|
||||
httpx.WriteTimeout,
|
||||
httpx.PoolTimeout,
|
||||
):
|
||||
assert is_connection_error(exc_cls("timeout")) is True
|
||||
|
||||
def test_httpx_network_errors(self):
|
||||
for exc_cls in (httpx.NetworkError, httpx.ProxyError, httpx.RequestError):
|
||||
assert is_connection_error(exc_cls("err")) is True
|
||||
|
||||
def test_builtin_network_exceptions(self):
|
||||
for exc in (ConnectionError("conn"), TimeoutError("timeout"), OSError("os")):
|
||||
assert is_connection_error(exc) is True
|
||||
|
||||
def test_non_network_exceptions_return_false(self):
|
||||
for exc in (
|
||||
ValueError("v"),
|
||||
TypeError("t"),
|
||||
RuntimeError("r"),
|
||||
BaseException(),
|
||||
):
|
||||
assert is_connection_error(exc) is False
|
||||
|
||||
def test_httpx_non_network_errors_return_false(self):
|
||||
request = httpx.Request("GET", "http://example.com")
|
||||
response = httpx.Response(200, request=request)
|
||||
for exc in (
|
||||
httpx.HTTPStatusError("err", request=request, response=response),
|
||||
httpx.InvalidURL("invalid"),
|
||||
httpx.CookieConflict("cookies"),
|
||||
):
|
||||
assert is_connection_error(exc) is False
|
||||
|
||||
def test_cause_chain_unwraps_to_network_error(self):
|
||||
inner = httpx.ConnectError("inner")
|
||||
outer = ValueError("wrapper")
|
||||
outer.__cause__ = inner
|
||||
assert is_connection_error(outer) is True
|
||||
|
||||
def test_cause_chain_no_network_error(self):
|
||||
inner = ValueError("inner")
|
||||
outer = RuntimeError("outer")
|
||||
outer.__cause__ = inner
|
||||
assert is_connection_error(outer) is False
|
||||
|
||||
def test_self_cause_does_not_infinite_loop(self):
|
||||
exc = ValueError("self")
|
||||
exc.__cause__ = exc
|
||||
assert is_connection_error(exc) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# log_connection_failure
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLogConnectionFailure:
|
||||
@patch("astrbot.core.utils.network_utils.logger.error")
|
||||
def test_with_explicit_proxy(self, mock_log_error: MagicMock):
|
||||
log_connection_failure("GPT", Exception("fail"), proxy="http://proxy:8080")
|
||||
mock_log_error.assert_called_once()
|
||||
msg = mock_log_error.call_args[0][0]
|
||||
assert "GPT" in msg
|
||||
assert "http://proxy:8080" in msg
|
||||
assert "fail" in msg
|
||||
|
||||
@patch("astrbot.core.utils.network_utils.logger.error")
|
||||
def test_without_proxy_falls_to_simple_message(self, mock_log_error: MagicMock):
|
||||
log_connection_failure("GPT", Exception("fail"))
|
||||
mock_log_error.assert_called_once()
|
||||
msg = mock_log_error.call_args[0][0]
|
||||
assert "网络连接失败" in msg
|
||||
assert "GPT" in msg
|
||||
|
||||
@patch("astrbot.core.utils.network_utils.logger.error")
|
||||
@patch.dict(os.environ, {"http_proxy": "http://env-proxy:3128"}, clear=False)
|
||||
def test_falls_back_to_env_http_proxy(
|
||||
self, mock_log_error: MagicMock
|
||||
):
|
||||
log_connection_failure("GPT", Exception("fail"))
|
||||
mock_log_error.assert_called_once()
|
||||
msg = mock_log_error.call_args[0][0]
|
||||
assert "http://env-proxy:3128" in msg
|
||||
|
||||
@patch("astrbot.core.utils.network_utils.logger.error")
|
||||
@patch.dict(os.environ, {"https_proxy": "https://secure-proxy:8443"}, clear=False)
|
||||
def test_falls_back_to_env_https_proxy(
|
||||
self, mock_log_error: MagicMock
|
||||
):
|
||||
log_connection_failure("GPT", Exception("fail"))
|
||||
mock_log_error.assert_called_once()
|
||||
msg = mock_log_error.call_args[0][0]
|
||||
assert "https://secure-proxy:8443" in msg
|
||||
|
||||
@patch("astrbot.core.utils.network_utils.logger.error")
|
||||
def test_empty_proxy_string_uses_no_proxy_message(
|
||||
self, mock_log_error: MagicMock
|
||||
):
|
||||
"""Empty-string proxy should not trigger the proxy-specific log line."""
|
||||
with patch.dict(os.environ, {"http_proxy": "", "https_proxy": ""}):
|
||||
log_connection_failure("GPT", Exception("fail"), proxy="")
|
||||
mock_log_error.assert_called_once()
|
||||
msg = mock_log_error.call_args[0][0]
|
||||
assert "代理" not in msg
|
||||
|
||||
@patch("astrbot.core.utils.network_utils.logger.error")
|
||||
def test_error_type_name_in_message(self, mock_log_error: MagicMock):
|
||||
log_connection_failure(
|
||||
"GPT", httpx.ConnectTimeout("upstream timed out")
|
||||
)
|
||||
mock_log_error.assert_called_once()
|
||||
msg = mock_log_error.call_args[0][0]
|
||||
assert "ConnectTimeout" in msg
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_proxy_client
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCreateProxyClientIntegration:
|
||||
"""Integration-style smoke tests against the real httpx.AsyncClient."""
|
||||
|
||||
def test_returns_async_client_without_proxy(self):
|
||||
client = create_proxy_client("Test")
|
||||
assert isinstance(client, httpx.AsyncClient)
|
||||
client.aclose()
|
||||
|
||||
def test_returns_async_client_with_proxy(self):
|
||||
client = create_proxy_client("Test", proxy="http://127.0.0.1:7890")
|
||||
assert isinstance(client, httpx.AsyncClient)
|
||||
client.aclose()
|
||||
|
||||
def test_with_custom_headers(self):
|
||||
headers = {"Authorization": "Bearer token"}
|
||||
client = create_proxy_client("Test", headers=headers)
|
||||
assert isinstance(client, httpx.AsyncClient)
|
||||
client.aclose()
|
||||
|
||||
def test_with_explicit_ssl_context(self):
|
||||
ctx = ssl.create_default_context()
|
||||
client = create_proxy_client("Test", verify=ctx)
|
||||
assert isinstance(client, httpx.AsyncClient)
|
||||
client.aclose()
|
||||
|
||||
def test_with_verify_disabled(self):
|
||||
client = create_proxy_client("Test", verify=False)
|
||||
assert isinstance(client, httpx.AsyncClient)
|
||||
client.aclose()
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
"""Unit tests for astrbot.core.pipeline.bootstrap module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.pipeline.bootstrap import (
|
||||
_BUILTIN_STAGE_MODULES,
|
||||
_EXPECTED_STAGE_NAMES,
|
||||
_builtin_stages_registered,
|
||||
ensure_builtin_stages_registered,
|
||||
)
|
||||
|
||||
|
||||
class TestEnsureBuiltinStagesRegistered:
|
||||
"""Tests for ensure_builtin_stages_registered()."""
|
||||
|
||||
def teardown_method(self):
|
||||
"""Reset global state after each test."""
|
||||
# Use import and direct assignment to reset
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_already_registered_flag_short_circuits(self):
|
||||
"""When _builtin_stages_registered is True, return immediately."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = True
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.bootstrap.import_module",
|
||||
) as mock_import:
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
mock_import.assert_not_called()
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_all_expected_stages_present_skips_import(self):
|
||||
"""When all expected stages are already in registered_stages, skip importing."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
|
||||
# Populate registered_stages with all expected stage names
|
||||
for name in _EXPECTED_STAGE_NAMES:
|
||||
cls = MagicMock()
|
||||
cls.__name__ = name
|
||||
bootstrap_mod.registered_stages.append(cls)
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.bootstrap.import_module",
|
||||
) as mock_import:
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
mock_import.assert_not_called()
|
||||
assert bootstrap_mod._builtin_stages_registered is True
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_registers_missing_stages(self):
|
||||
"""When stages are missing, import built-in modules."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
bootstrap_mod.registered_stages.clear()
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.bootstrap.import_module",
|
||||
) as mock_import:
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
expected_calls = [call(mod) for mod in _BUILTIN_STAGE_MODULES]
|
||||
mock_import.assert_has_calls(expected_calls, any_order=True)
|
||||
assert mock_import.call_count == len(_BUILTIN_STAGE_MODULES)
|
||||
assert bootstrap_mod._builtin_stages_registered is True
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_idempotent(self):
|
||||
"""Calling ensure_builtin_stages_registered twice only imports once."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
bootstrap_mod.registered_stages.clear()
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.bootstrap.import_module",
|
||||
) as mock_import:
|
||||
ensure_builtin_stages_registered()
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
# Should only import on first call
|
||||
mock_import.assert_called_once()
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_global_flag_set_after_import(self):
|
||||
"""Verify the global flag is set after a full registration."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
bootstrap_mod.registered_stages.clear()
|
||||
|
||||
with patch("astrbot.core.pipeline.bootstrap.import_module"):
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
assert bootstrap_mod._builtin_stages_registered is True
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_already_registered_flag_persists(self):
|
||||
"""Verify that once set, the flag persists across calls."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
bootstrap_mod.registered_stages.clear()
|
||||
|
||||
with patch("astrbot.core.pipeline.bootstrap.import_module"):
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
# Call again with import_module mocked to raise if called
|
||||
with patch(
|
||||
"astrbot.core.pipeline.bootstrap.import_module",
|
||||
) as mock_import:
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
mock_import.assert_not_called()
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_partial_stages_present_still_imports(self):
|
||||
"""When only some expected stages are present, still import all modules."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
bootstrap_mod.registered_stages.clear()
|
||||
|
||||
# Add only one of the expected stages
|
||||
cls = MagicMock()
|
||||
cls.__name__ = "ProcessStage"
|
||||
bootstrap_mod.registered_stages.append(cls)
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.bootstrap.import_module",
|
||||
) as mock_import:
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
expected_calls = [call(mod) for mod in _BUILTIN_STAGE_MODULES]
|
||||
mock_import.assert_has_calls(expected_calls, any_order=True)
|
||||
assert mock_import.call_count == len(_BUILTIN_STAGE_MODULES)
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_stage_names_check_exact(self):
|
||||
"""Verify the check uses __name__ comparison, not identity."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
bootstrap_mod.registered_stages.clear()
|
||||
|
||||
# Add stages with correct names
|
||||
for name in _EXPECTED_STAGE_NAMES:
|
||||
cls = MagicMock()
|
||||
cls.__name__ = name
|
||||
bootstrap_mod.registered_stages.append(cls)
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.bootstrap.import_module",
|
||||
) as mock_import:
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
mock_import.assert_not_called()
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_expected_stage_names_are_correct(self):
|
||||
"""Verify _EXPECTED_STAGE_NAMES matches the expected set."""
|
||||
expected = {
|
||||
"WakingCheckStage",
|
||||
"WhitelistCheckStage",
|
||||
"SessionStatusCheckStage",
|
||||
"RateLimitStage",
|
||||
"ContentSafetyCheckStage",
|
||||
"PreProcessStage",
|
||||
"ProcessStage",
|
||||
"ResultDecorateStage",
|
||||
"RespondStage",
|
||||
}
|
||||
assert _EXPECTED_STAGE_NAMES == expected
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_builtin_stage_modules_count(self):
|
||||
"""Verify the number of builtin stage modules matches expected."""
|
||||
assert len(_BUILTIN_STAGE_MODULES) == 9
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_all_builtin_modules_under_pipeline(self):
|
||||
"""Verify all builtin modules are under astrbot.core.pipeline."""
|
||||
for mod in _BUILTIN_STAGE_MODULES:
|
||||
assert mod.startswith("astrbot.core.pipeline.")
|
||||
assert mod.endswith(".stage")
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_module_paths_alignment(self):
|
||||
"""Verify each expected stage name has a corresponding module path."""
|
||||
from astrbot.core.pipeline.bootstrap import _BUILTIN_STAGE_MODULES, _EXPECTED_STAGE_NAMES
|
||||
|
||||
# Derive expected names from module paths
|
||||
derived_names = set()
|
||||
for mod_path in _BUILTIN_STAGE_MODULES:
|
||||
parts = mod_path.split(".")
|
||||
# For modules like astrbot.core.pipeline.process_stage.stage
|
||||
# the stage name is found in the penultimate segment
|
||||
if parts[-2] in ("process_stage", "respond"):
|
||||
# Special cases: ProcessStage, RespondStage
|
||||
if parts[-2] == "process_stage":
|
||||
derived_names.add("ProcessStage")
|
||||
elif parts[-2] == "respond":
|
||||
derived_names.add("RespondStage")
|
||||
else:
|
||||
derived_names.add(f"{parts[-2].title().replace('_', '')}Stage")
|
||||
else:
|
||||
name = parts[-2].replace("_", " ").title().replace(" ", "")
|
||||
name += "Stage"
|
||||
derived_names.add(name)
|
||||
|
||||
assert _EXPECTED_STAGE_NAMES == derived_names
|
||||
|
||||
@patch("astrbot.core.pipeline.bootstrap.registered_stages", new=[])
|
||||
def test_real_registration_smoke(self):
|
||||
"""Smoke test: calling ensure_builtin_stages_registered with actual modules."""
|
||||
import astrbot.core.pipeline.bootstrap as bootstrap_mod
|
||||
|
||||
bootstrap_mod._builtin_stages_registered = False
|
||||
bootstrap_mod.registered_stages.clear()
|
||||
|
||||
# This should import the actual modules
|
||||
ensure_builtin_stages_registered()
|
||||
|
||||
stage_names = {cls.__name__ for cls in bootstrap_mod.registered_stages}
|
||||
assert _EXPECTED_STAGE_NAMES.issubset(stage_names)
|
||||
@@ -0,0 +1,346 @@
|
||||
"""Unit tests for astrbot.core.pipeline.process_stage.stage.ProcessStage."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.pipeline.process_stage.stage import ProcessStage
|
||||
from astrbot.core.provider.entities import ProviderRequest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_context():
|
||||
"""Create a mock PipelineContext."""
|
||||
ctx = MagicMock()
|
||||
ctx.astrbot_config = {
|
||||
"provider_settings": {"enable": True, "wake_prefix": ""},
|
||||
"wake_prefix": ["bot"],
|
||||
}
|
||||
ctx.plugin_manager = MagicMock()
|
||||
ctx.plugin_manager.context = MagicMock()
|
||||
return ctx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_event():
|
||||
"""Create a mock AstrMessageEvent."""
|
||||
event = MagicMock()
|
||||
event.get_extra.return_value = None
|
||||
event.is_stopped.return_value = False
|
||||
event._has_send_oper = False
|
||||
event.is_at_or_wake_command = True
|
||||
event.call_llm = True
|
||||
event.set_extra = MagicMock()
|
||||
return event
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stage(mock_context):
|
||||
"""Create a ProcessStage with mocked sub-stages."""
|
||||
stage = ProcessStage()
|
||||
|
||||
stage.agent_sub_stage = MagicMock()
|
||||
stage.agent_sub_stage.process = AsyncMock()
|
||||
|
||||
stage.star_request_sub_stage = MagicMock()
|
||||
stage.star_request_sub_stage.process = AsyncMock()
|
||||
|
||||
stage.ctx = mock_context
|
||||
stage.config = mock_context.astrbot_config
|
||||
stage.plugin_manager = mock_context.plugin_manager
|
||||
stage.sdk_plugin_bridge = None
|
||||
return stage
|
||||
|
||||
|
||||
class TestProcessStageInitialize:
|
||||
"""Tests for ProcessStage.initialize()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_sets_attributes(self, mock_context):
|
||||
"""Verify initialize sets ctx, config, plugin_manager."""
|
||||
stage = ProcessStage()
|
||||
with (
|
||||
patch.object(stage, "agent_sub_stage") as mock_agent,
|
||||
patch.object(stage, "star_request_sub_stage") as mock_star,
|
||||
):
|
||||
await stage.initialize(mock_context)
|
||||
|
||||
assert stage.ctx is mock_context
|
||||
assert stage.config is mock_context.astrbot_config
|
||||
assert stage.plugin_manager is mock_context.plugin_manager
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_creates_sub_stages(self, mock_context):
|
||||
"""Verify initialize creates and initializes sub-stages."""
|
||||
stage = ProcessStage()
|
||||
await stage.initialize(mock_context)
|
||||
|
||||
assert hasattr(stage, "agent_sub_stage")
|
||||
assert hasattr(stage, "star_request_sub_stage")
|
||||
assert stage.agent_sub_stage is not None
|
||||
assert stage.star_request_sub_stage is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_sdk_plugin_bridge_present(self, mock_context):
|
||||
"""Verify sdk_plugin_bridge is set when plugin_manager has it."""
|
||||
mock_context.plugin_manager.context.sdk_plugin_bridge = MagicMock()
|
||||
stage = ProcessStage()
|
||||
await stage.initialize(mock_context)
|
||||
|
||||
assert stage.sdk_plugin_bridge is mock_context.plugin_manager.context.sdk_plugin_bridge
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_sdk_plugin_bridge_absent(self, mock_context):
|
||||
"""Verify sdk_plugin_bridge is None when plugin_manager has no context attr."""
|
||||
mock_context.plugin_manager = MagicMock(spec=[]) # no context attribute
|
||||
stage = ProcessStage()
|
||||
await stage.initialize(mock_context)
|
||||
|
||||
assert stage.sdk_plugin_bridge is None
|
||||
|
||||
|
||||
class TestProcessStageProcess:
|
||||
"""Tests for ProcessStage.process()."""
|
||||
|
||||
async def _collect(self, async_gen):
|
||||
"""Helper to collect all items from an async generator."""
|
||||
results = []
|
||||
async for item in async_gen:
|
||||
results.append(item)
|
||||
return results
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_no_activated_handlers_and_no_sdk_bridge(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""When activated_handlers is None and no sdk bridge, skip handler path."""
|
||||
mock_event.get_extra.return_value = None
|
||||
stage.sdk_plugin_bridge = None
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
assert results == [] # provider enabled, should still reach agent_sub_stage
|
||||
stage.agent_sub_stage.process.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_activated_handlers_with_provider_request(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""When star_request yields a ProviderRequest, agent_sub_stage is called."""
|
||||
pr = ProviderRequest(prompt="hello")
|
||||
|
||||
async def star_gen(_event):
|
||||
yield pr
|
||||
|
||||
stage.star_request_sub_stage.process = star_gen
|
||||
|
||||
agent_called = False
|
||||
|
||||
async def agent_gen(_event):
|
||||
nonlocal agent_called
|
||||
agent_called = True
|
||||
yield None
|
||||
|
||||
stage.agent_sub_stage.process = agent_gen
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
assert agent_called is True
|
||||
mock_event.set_extra.assert_any_call("provider_request", pr)
|
||||
assert len(results) >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_activated_handlers_with_non_provider_request(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""When star_request yields a non-ProviderRequest, yield directly."""
|
||||
async def star_gen(_event):
|
||||
yield "some_other_result"
|
||||
|
||||
stage.star_request_sub_stage.process = star_gen
|
||||
stage.sdk_plugin_bridge = None
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
assert len(results) >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_activated_handlers_provider_request_empty_agent(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""When agent_sub_stage yields nothing, should still yield once."""
|
||||
pr = ProviderRequest(prompt="hi")
|
||||
|
||||
async def star_gen(_event):
|
||||
yield pr
|
||||
|
||||
stage.star_request_sub_stage.process = star_gen
|
||||
|
||||
async def empty_agent_gen(_event):
|
||||
if False:
|
||||
yield None
|
||||
|
||||
stage.agent_sub_stage.process = empty_agent_gen
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
assert len(results) >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_sdk_plugin_bridge_sent_message(self, stage, mock_event):
|
||||
"""When sdk_plugin_bridge.dispatch_message returns sent_message=True."""
|
||||
mock_bridge = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.sent_message = True
|
||||
mock_result.stopped = False
|
||||
mock_bridge.dispatch_message = AsyncMock(return_value=mock_result)
|
||||
stage.sdk_plugin_bridge = mock_bridge
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
# Should have yielded due to sent_message
|
||||
assert len(results) >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_sdk_plugin_bridge_stopped(self, stage, mock_event):
|
||||
"""When sdk_plugin_bridge.dispatch_message returns stopped=True."""
|
||||
mock_bridge = MagicMock()
|
||||
mock_result = MagicMock()
|
||||
mock_result.sent_message = False
|
||||
mock_result.stopped = True
|
||||
mock_bridge.dispatch_message = AsyncMock(return_value=mock_result)
|
||||
stage.sdk_plugin_bridge = mock_bridge
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
# Should have yielded due to stopped
|
||||
assert len(results) >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_sdk_plugin_bridge_none_and_event_has_send_oper(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""When sdk bridge is None and _has_send_oper is True, skip LLM call."""
|
||||
stage.sdk_plugin_bridge = None
|
||||
mock_event._has_send_oper = True
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
# Should NOT call agent_sub_stage for LLM
|
||||
assert results == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_provider_disabled(self, stage, mock_event):
|
||||
"""When provider_settings enable is False, return early."""
|
||||
stage.ctx.astrbot_config["provider_settings"]["enable"] = False
|
||||
stage.sdk_plugin_bridge = None
|
||||
mock_event.get_extra.return_value = None
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
assert results == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_llm_triggered(self, stage, mock_event):
|
||||
"""When all conditions met, agent_sub_stage is called for LLM."""
|
||||
stage.sdk_plugin_bridge = None
|
||||
agent_called = False
|
||||
|
||||
async def agent_gen(_event):
|
||||
nonlocal agent_called
|
||||
agent_called = True
|
||||
yield None
|
||||
|
||||
stage.agent_sub_stage.process = agent_gen
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
assert agent_called is True
|
||||
assert len(results) >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_llm_skipped_not_at_wake(self, stage, mock_event):
|
||||
"""When is_at_or_wake_command is False, skip LLM call."""
|
||||
stage.sdk_plugin_bridge = None
|
||||
mock_event.is_at_or_wake_command = False
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
assert results == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_llm_skipped_event_stopped_with_result(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""When event is stopped and effective_result exists, skip LLM call."""
|
||||
stage.sdk_plugin_bridge = None
|
||||
mock_event.is_stopped.return_value = True
|
||||
|
||||
mock_bridge = MagicMock()
|
||||
mock_bridge.get_effective_should_call_llm.return_value = True
|
||||
mock_bridge.get_effective_result.return_value = "some_result"
|
||||
stage.sdk_plugin_bridge = mock_bridge
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
# Event is stopped with result, skip LLM
|
||||
assert results == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_sdk_bridge_with_effective_methods(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""When sdk_plugin_bridge has get_effective_* methods, they are used."""
|
||||
mock_bridge = MagicMock()
|
||||
mock_bridge.get_effective_should_call_llm.return_value = False
|
||||
stage.sdk_plugin_bridge = mock_bridge
|
||||
|
||||
await self._collect(stage.process(mock_event))
|
||||
|
||||
mock_bridge.get_effective_should_call_llm.assert_called_once_with(
|
||||
mock_event,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_with_activated_handlers_and_stopped(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""When event is stopped after star_request processing, stop propagation."""
|
||||
async def star_gen(_event):
|
||||
yield ProviderRequest(prompt="test")
|
||||
|
||||
stage.star_request_sub_stage.process = star_gen
|
||||
|
||||
async def agent_gen(_event):
|
||||
yield None
|
||||
|
||||
stage.agent_sub_stage.process = agent_gen
|
||||
|
||||
mock_event.is_stopped.return_value = True
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
assert len(results) >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_event_call_llm_false(self, stage, mock_event):
|
||||
"""When event.call_llm is False and no sdk bridge, should_call_llm is True (inverted)."""
|
||||
stage.sdk_plugin_bridge = None
|
||||
mock_event.call_llm = False
|
||||
|
||||
agent_called = False
|
||||
|
||||
async def agent_gen(_event):
|
||||
nonlocal agent_called
|
||||
agent_called = True
|
||||
yield None
|
||||
|
||||
stage.agent_sub_stage.process = agent_gen
|
||||
|
||||
results = await self._collect(stage.process(mock_event))
|
||||
|
||||
# event.call_llm is False, and no sdk_bridge -> should_call_llm = not False = True
|
||||
assert agent_called is True
|
||||
assert len(results) >= 1
|
||||
@@ -0,0 +1,530 @@
|
||||
"""Unit tests for astrbot.core.pipeline.respond.stage.RespondStage."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.message.components import (
|
||||
At,
|
||||
ComponentType,
|
||||
Face,
|
||||
File,
|
||||
Forward,
|
||||
Image,
|
||||
Plain,
|
||||
Record,
|
||||
Reply,
|
||||
)
|
||||
from astrbot.core.message.message_event_result import (
|
||||
MessageEventResult,
|
||||
ResultContentType,
|
||||
)
|
||||
from astrbot.core.pipeline.respond.stage import RespondStage
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_config():
|
||||
"""Create a mock AstrBotConfig for RespondStage initialization."""
|
||||
return {
|
||||
"platform_settings": {
|
||||
"reply_with_mention": True,
|
||||
"reply_with_quote": False,
|
||||
"segmented_reply": {
|
||||
"enable": False,
|
||||
"only_llm_result": False,
|
||||
"interval_method": "random",
|
||||
"log_base": 2,
|
||||
"interval": "1.5, 3.5",
|
||||
},
|
||||
"path_mapping": [],
|
||||
},
|
||||
"provider_settings": {
|
||||
"enable": True,
|
||||
"unsupported_streaming_strategy": "realtime_segmenting",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_context(mock_config):
|
||||
"""Create a mock PipelineContext."""
|
||||
ctx = MagicMock()
|
||||
ctx.astrbot_config = mock_config
|
||||
return ctx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_event():
|
||||
"""Create a mock AstrMessageEvent."""
|
||||
event = MagicMock()
|
||||
event.get_extra.return_value = False
|
||||
event.get_platform_name.return_value = "qq"
|
||||
event.get_sender_name.return_value = "test_user"
|
||||
event.get_sender_id.return_value = "12345"
|
||||
event.get_platform_id.return_value = "platform_1"
|
||||
return event
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stage(mock_context):
|
||||
"""Create an initialized RespondStage."""
|
||||
stage = RespondStage()
|
||||
return stage
|
||||
|
||||
|
||||
class TestRespondStageInitialize:
|
||||
"""Tests for RespondStage.initialize()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_sets_attributes(self, stage, mock_context):
|
||||
"""Verify initialize reads config and sets attributes."""
|
||||
await stage.initialize(mock_context)
|
||||
|
||||
assert stage.ctx is mock_context
|
||||
assert stage.config is mock_context.astrbot_config
|
||||
assert stage.reply_with_mention is True
|
||||
assert stage.reply_with_quote is False
|
||||
assert stage.enable_seg is False
|
||||
assert stage.only_llm_result is False
|
||||
assert stage.interval_method == "random"
|
||||
assert stage.log_base == 2
|
||||
assert stage.interval == [1.5, 3.5]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_segmented_enabled(self, mock_context):
|
||||
"""Verify segmented reply config is parsed when enabled."""
|
||||
mock_context.astrbot_config["platform_settings"]["segmented_reply"][
|
||||
"enable"
|
||||
] = True
|
||||
stage = RespondStage()
|
||||
await stage.initialize(mock_context)
|
||||
|
||||
assert stage.enable_seg is True
|
||||
assert stage.interval == [1.5, 3.5]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_segmented_invalid_interval(self, mock_context):
|
||||
"""Verify invalid interval string falls back gracefully."""
|
||||
mock_context.astrbot_config["platform_settings"]["segmented_reply"][
|
||||
"enable"
|
||||
] = True
|
||||
mock_context.astrbot_config["platform_settings"]["segmented_reply"][
|
||||
"interval"
|
||||
] = "not_a_number"
|
||||
stage = RespondStage()
|
||||
await stage.initialize(mock_context)
|
||||
|
||||
# Should fall back to [1.5, 3.5]
|
||||
assert stage.interval == [1.5, 3.5]
|
||||
|
||||
|
||||
class TestRespondStageWordCnt:
|
||||
"""Tests for RespondStage._word_cnt()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_word_cnt_ascii(self, stage):
|
||||
"""Verify ASCII text is split and counted."""
|
||||
count = await stage._word_cnt("hello world foo bar")
|
||||
assert count == 4
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_word_cnt_ascii_single(self, stage):
|
||||
"""Verify single ASCII word."""
|
||||
count = await stage._word_cnt("hello")
|
||||
assert count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_word_cnt_non_ascii(self, stage):
|
||||
"""Verify non-ASCII text counts alphanumeric characters."""
|
||||
count = await stage._word_cnt("你好世界")
|
||||
assert count == 4
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_word_cnt_mixed(self, stage):
|
||||
"""Verify mixed text counts all alnum characters."""
|
||||
count = await stage._word_cnt("hello你好world世界")
|
||||
assert count == 18 # hello(5) + 你好(4) + world(5) + 世界(4) = 18
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_word_cnt_empty(self, stage):
|
||||
"""Verify empty string returns 0."""
|
||||
count = await stage._word_cnt("")
|
||||
assert count == 0
|
||||
|
||||
|
||||
class TestRespondStageCalcCompInterval:
|
||||
"""Tests for RespondStage._calc_comp_interval()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_calc_interval_log_plain(self, stage):
|
||||
"""Verify log interval for Plain component."""
|
||||
stage.interval_method = "log"
|
||||
stage.log_base = 2
|
||||
plain = Plain(text="hello world")
|
||||
|
||||
with patch("astrbot.core.pipeline.respond.stage.random.uniform") as mock_uniform:
|
||||
mock_uniform.return_value = 2.0
|
||||
interval = await stage._calc_comp_interval(plain)
|
||||
|
||||
assert interval == 2.0
|
||||
mock_uniform.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_calc_interval_log_non_plain(self, stage):
|
||||
"""Verify log interval for non-Plain component uses 1-1.75 range."""
|
||||
stage.interval_method = "log"
|
||||
image = Image(file="/path/to/img.jpg")
|
||||
|
||||
with patch("astrbot.core.pipeline.respond.stage.random.uniform") as mock_uniform:
|
||||
mock_uniform.return_value = 1.5
|
||||
interval = await stage._calc_comp_interval(image)
|
||||
|
||||
assert interval == 1.5
|
||||
mock_uniform.assert_called_once_with(1, 1.75)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_calc_interval_random(self, stage):
|
||||
"""Verify random interval uses configured interval range."""
|
||||
stage.interval_method = "random"
|
||||
stage.interval = [2.0, 4.0]
|
||||
plain = Plain(text="test")
|
||||
|
||||
with patch("astrbot.core.pipeline.respond.stage.random.uniform") as mock_uniform:
|
||||
mock_uniform.return_value = 3.0
|
||||
interval = await stage._calc_comp_interval(plain)
|
||||
|
||||
assert interval == 3.0
|
||||
mock_uniform.assert_called_once_with(2.0, 4.0)
|
||||
|
||||
|
||||
class TestRespondStageHasMeaningfulContent:
|
||||
"""Tests for RespondStage._has_meaningful_content()."""
|
||||
|
||||
def test_plain_with_text(self, stage):
|
||||
"""Verify Plain with text returns True."""
|
||||
comp = Plain(text="hello")
|
||||
assert stage._has_meaningful_content(comp) is True
|
||||
|
||||
def test_plain_whitespace_text(self, stage):
|
||||
"""Verify Plain with whitespace returns False."""
|
||||
comp = Plain(text=" ")
|
||||
assert stage._has_meaningful_content(comp) is False
|
||||
|
||||
def test_image_with_url(self, stage):
|
||||
"""Verify Image with url returns True."""
|
||||
comp = Image(file="http://example.com/img.jpg")
|
||||
assert stage._has_meaningful_content(comp) is True
|
||||
|
||||
def test_image_with_file_id(self, stage):
|
||||
"""Verify Image with file_id returns True."""
|
||||
comp = Image(file_id="abc123")
|
||||
assert stage._has_meaningful_content(comp) is True
|
||||
|
||||
def test_image_empty(self, stage):
|
||||
"""Verify Image without url or file_id returns False."""
|
||||
comp = Image()
|
||||
assert stage._has_meaningful_content(comp) is False
|
||||
|
||||
def test_face_with_id(self, stage):
|
||||
"""Verify Face with id returns True."""
|
||||
comp = Face(id=123)
|
||||
assert stage._has_meaningful_content(comp) is True
|
||||
|
||||
def test_face_no_id(self, stage):
|
||||
"""Verify Face without id returns False."""
|
||||
comp = Face(id=None)
|
||||
assert stage._has_meaningful_content(comp) is False
|
||||
|
||||
def test_at_with_qq(self, stage):
|
||||
"""Verify At with qq returns True."""
|
||||
comp = At(qq="12345")
|
||||
assert stage._has_meaningful_content(comp) is True
|
||||
|
||||
def test_at_no_qq(self, stage):
|
||||
"""Verify At without qq returns False."""
|
||||
comp = At(qq=None)
|
||||
assert stage._has_meaningful_content(comp) is False
|
||||
|
||||
def test_reply_with_id(self, stage):
|
||||
"""Verify Reply with id returns True."""
|
||||
comp = Reply(id="abc", sender_id="user1")
|
||||
assert stage._has_meaningful_content(comp) is True
|
||||
|
||||
def test_reply_no_id(self, stage):
|
||||
"""Verify Reply without id returns False."""
|
||||
comp = Reply(id=None, sender_id="user1")
|
||||
assert stage._has_meaningful_content(comp) is False
|
||||
|
||||
def test_forward_with_id(self, stage):
|
||||
"""Verify Forward with id returns True."""
|
||||
comp = Forward(id="abc123")
|
||||
assert stage._has_meaningful_content(comp) is True
|
||||
|
||||
|
||||
class TestRespondStageIsEmptyMessageChain:
|
||||
"""Tests for RespondStage._is_empty_message_chain()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_list(self, stage):
|
||||
"""Verify empty list returns True."""
|
||||
assert await stage._is_empty_message_chain([]) is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chain_with_meaningful_content(self, stage):
|
||||
"""Verify chain with valid content returns False."""
|
||||
chain = [Plain(text="hello")]
|
||||
assert await stage._is_empty_message_chain(chain) is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chain_all_empty(self, stage):
|
||||
"""Verify chain with all empty components returns True."""
|
||||
chain = [Plain(text="")]
|
||||
assert await stage._is_empty_message_chain(chain) is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chain_mixed_empty_and_valid(self, stage):
|
||||
"""Verify chain with mix of empty and valid returns False."""
|
||||
chain = [Plain(text=""), Plain(text="hello")]
|
||||
assert await stage._is_empty_message_chain(chain) is False
|
||||
|
||||
|
||||
class TestRespondStageIsSegReplyRequired:
|
||||
"""Tests for RespondStage.is_seg_reply_required()."""
|
||||
|
||||
def test_seg_disabled(self, stage):
|
||||
"""Verify returns False when segmented reply is disabled."""
|
||||
stage.enable_seg = False
|
||||
event = MagicMock()
|
||||
assert stage.is_seg_reply_required(event) is False
|
||||
|
||||
def test_seg_no_result(self, stage):
|
||||
"""Verify returns False when event has no result."""
|
||||
stage.enable_seg = True
|
||||
event = MagicMock()
|
||||
event.get_result.return_value = None
|
||||
assert stage.is_seg_reply_required(event) is False
|
||||
|
||||
def test_seg_only_llm_not_model(self, stage):
|
||||
"""Verify returns False when only_llm_result is True and result is not model."""
|
||||
stage.enable_seg = True
|
||||
stage.only_llm_result = True
|
||||
event = MagicMock()
|
||||
result = MagicMock()
|
||||
result.is_model_result.return_value = False
|
||||
event.get_result.return_value = result
|
||||
assert stage.is_seg_reply_required(event) is False
|
||||
|
||||
def test_seg_excluded_platform(self, stage):
|
||||
"""Verify returns False for excluded platforms."""
|
||||
stage.enable_seg = True
|
||||
stage.only_llm_result = False
|
||||
event = MagicMock()
|
||||
result = MagicMock()
|
||||
result.is_model_result.return_value = True
|
||||
event.get_result.return_value = result
|
||||
event.get_platform_name.return_value = "qq_official"
|
||||
assert stage.is_seg_reply_required(event) is False
|
||||
|
||||
def test_seg_all_conditions_met(self, stage):
|
||||
"""Verify returns True when all conditions are met."""
|
||||
stage.enable_seg = True
|
||||
stage.only_llm_result = False
|
||||
event = MagicMock()
|
||||
result = MagicMock()
|
||||
result.is_model_result.return_value = True
|
||||
event.get_result.return_value = result
|
||||
event.get_platform_name.return_value = "qq"
|
||||
assert stage.is_seg_reply_required(event) is True
|
||||
|
||||
|
||||
class TestRespondStageExtractComp:
|
||||
"""Tests for RespondStage._extract_comp()."""
|
||||
|
||||
def test_extract_with_modify(self, stage):
|
||||
"""Verify extraction removes extracted types from original list."""
|
||||
raw_chain = [
|
||||
Plain(text="hello"),
|
||||
At(qq="12345"),
|
||||
Plain(text="world"),
|
||||
Reply(id="r1", sender_id="u1"),
|
||||
]
|
||||
extracted = stage._extract_comp(
|
||||
raw_chain,
|
||||
{ComponentType.At, ComponentType.Reply},
|
||||
modify_raw_chain=True,
|
||||
)
|
||||
|
||||
assert len(extracted) == 2
|
||||
assert all(c.type in {ComponentType.At, ComponentType.Reply} for c in extracted)
|
||||
assert len(raw_chain) == 2
|
||||
assert all(c.type == ComponentType.Plain for c in raw_chain)
|
||||
|
||||
def test_extract_without_modify(self, stage):
|
||||
"""Verify extraction does not modify original chain."""
|
||||
raw_chain = [
|
||||
Plain(text="hello"),
|
||||
At(qq="12345"),
|
||||
]
|
||||
original_len = len(raw_chain)
|
||||
extracted = stage._extract_comp(
|
||||
raw_chain,
|
||||
{ComponentType.At},
|
||||
modify_raw_chain=False,
|
||||
)
|
||||
|
||||
assert len(extracted) == 1
|
||||
assert extracted[0].type == ComponentType.At
|
||||
assert len(raw_chain) == original_len # unchanged
|
||||
|
||||
def test_extract_no_match(self, stage):
|
||||
"""Verify extraction returns empty when no types match."""
|
||||
raw_chain = [Plain(text="hello")]
|
||||
extracted = stage._extract_comp(
|
||||
raw_chain,
|
||||
{ComponentType.Image},
|
||||
modify_raw_chain=False,
|
||||
)
|
||||
assert extracted == []
|
||||
|
||||
|
||||
class TestRespondStageProcess:
|
||||
"""Tests for RespondStage.process()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_result_none(self, stage, mock_event):
|
||||
"""Verify process returns when result is None."""
|
||||
mock_event.get_result.return_value = None
|
||||
await stage.process(mock_event)
|
||||
# Should return without sending
|
||||
mock_event.send.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_streaming_finished(self, stage, mock_event):
|
||||
"""Verify process returns when streaming is finished."""
|
||||
result = MagicMock()
|
||||
result.result_content_type = ResultContentType.STREAMING_FINISH
|
||||
mock_event.get_result.return_value = result
|
||||
await stage.process(mock_event)
|
||||
|
||||
mock_event.set_extra.assert_called_once_with("_streaming_finished", True)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_streaming_finish_prevents_duplicate_send(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""Verify prevent duplicate send after streaming finish."""
|
||||
mock_event.get_extra.return_value = True # _streaming_finished already True
|
||||
result = MagicMock()
|
||||
result.result_content_type = ResultContentType.GENERAL_RESULT
|
||||
mock_event.get_result.return_value = result
|
||||
await stage.process(mock_event)
|
||||
|
||||
mock_event.send.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_streaming_result(self, stage, mock_event):
|
||||
"""Verify STREAMING_RESULT is delivered directly to event.send_streaming."""
|
||||
async def dummy_async_stream():
|
||||
yield "chunk1"
|
||||
|
||||
result = MagicMock()
|
||||
result.result_content_type = ResultContentType.STREAMING_RESULT
|
||||
result.async_stream = dummy_async_stream()
|
||||
mock_event.get_result.return_value = result
|
||||
|
||||
await stage.process(mock_event)
|
||||
|
||||
mock_event.send_streaming.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_normal_chain(self, stage, mock_event):
|
||||
"""Verify normal message chain is sent via event.send."""
|
||||
result = MessageEventResult()
|
||||
result.chain = [Plain(text="hello")]
|
||||
mock_event.get_result.return_value = result
|
||||
|
||||
await stage.process(mock_event)
|
||||
|
||||
mock_event.send.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_empty_chain(self, stage, mock_event):
|
||||
"""Verify empty chain skips sending."""
|
||||
result = MessageEventResult()
|
||||
result.chain = [Plain(text="")]
|
||||
mock_event.get_result.return_value = result
|
||||
|
||||
await stage.process(mock_event)
|
||||
|
||||
mock_event.send.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_chain_with_path_mapping(self, stage, mock_event):
|
||||
"""Verify path mapping is applied to File components."""
|
||||
stage.platform_settings = {
|
||||
"path_mapping": [["/old/path", "/new/path"]],
|
||||
}
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.respond.stage.path_Mapping",
|
||||
return_value="/new/path/file.txt",
|
||||
):
|
||||
result = MessageEventResult()
|
||||
result.chain = [File(file="/old/path/file.txt")]
|
||||
mock_event.get_result.return_value = result
|
||||
|
||||
await stage.process(mock_event)
|
||||
|
||||
mock_event.send.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_chain_with_record_forced_separate(
|
||||
self, stage, mock_event,
|
||||
):
|
||||
"""Verify Record components are sent separately."""
|
||||
result = MessageEventResult()
|
||||
result.chain = [Record(file="audio.mp3"), Plain(text="hello")]
|
||||
mock_event.get_result.return_value = result
|
||||
|
||||
await stage.process(mock_event)
|
||||
|
||||
# Record should be sent separately, then the rest
|
||||
assert mock_event.send.call_count >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_segmented_reply(self, stage, mock_event):
|
||||
"""Verify segmented reply sends each component with delay."""
|
||||
stage.enable_seg = True
|
||||
stage.only_llm_result = False
|
||||
stage.interval_method = "random"
|
||||
stage.interval = [0.1, 0.2]
|
||||
|
||||
result = MessageEventResult()
|
||||
result.chain = [Plain(text="part1"), Plain(text="part2")]
|
||||
mock_event.get_result.return_value = result
|
||||
mock_event.get_platform_name.return_value = "qq"
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.respond.stage.asyncio.sleep",
|
||||
AsyncMock(),
|
||||
):
|
||||
await stage.process(mock_event)
|
||||
|
||||
assert mock_event.send.call_count >= 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_event_hook_called(self, stage, mock_event):
|
||||
"""Verify call_event_hook is invoked after sending."""
|
||||
result = MessageEventResult()
|
||||
result.chain = [Plain(text="hello")]
|
||||
mock_event.get_result.return_value = result
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.respond.stage.call_event_hook",
|
||||
AsyncMock(return_value=False),
|
||||
):
|
||||
await stage.process(mock_event)
|
||||
|
||||
mock_event.clear_result.assert_called_once()
|
||||
@@ -0,0 +1,400 @@
|
||||
"""Unit tests for astrbot.core.pipeline.scheduler.PipelineScheduler."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.pipeline.scheduler import PipelineScheduler
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_context():
|
||||
"""Create a mock PipelineContext."""
|
||||
ctx = MagicMock()
|
||||
ctx.astrbot_config = {"provider_settings": {"enable": True}}
|
||||
ctx.plugin_manager = MagicMock()
|
||||
ctx.plugin_manager.context = MagicMock(spec_set=[])
|
||||
return ctx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_stage_cls():
|
||||
"""Create a mock stage class that can be instantiated."""
|
||||
cls = MagicMock()
|
||||
instance = MagicMock()
|
||||
instance.initialize = AsyncMock()
|
||||
instance.process = AsyncMock()
|
||||
cls.return_value = instance
|
||||
cls.__name__ = "MockStage"
|
||||
return cls
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Init
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPipelineSchedulerInit:
|
||||
"""Tests for PipelineScheduler.__init__()."""
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.STAGES_ORDER", ["MockStage"])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
def test_init_calls_ensure_and_sets_context(
|
||||
self, mock_ensure, mock_stages_order, mock_registered, mock_context,
|
||||
):
|
||||
"""Verify __init__ calls ensure_builtin_stages_registered and sets context."""
|
||||
mock_stage = MagicMock()
|
||||
mock_stage.__name__ = "MockStage"
|
||||
mock_registered.append(mock_stage)
|
||||
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
|
||||
mock_ensure.assert_called_once()
|
||||
assert scheduler.ctx is mock_context
|
||||
assert scheduler.stages == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Initialize
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPipelineSchedulerInitialize:
|
||||
"""Tests for PipelineScheduler.initialize()."""
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_creates_stage_instances(
|
||||
self, mock_ensure, mock_registered, mock_context, mock_stage_cls,
|
||||
):
|
||||
"""Verify initialize creates and initializes all registered stage instances."""
|
||||
mock_registered.append(mock_stage_cls)
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
|
||||
await scheduler.initialize()
|
||||
|
||||
assert len(scheduler.stages) == 1
|
||||
assert scheduler.stages[0] is mock_stage_cls.return_value
|
||||
mock_stage_cls.return_value.initialize.assert_awaited_once_with(mock_context)
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_multiple_stages(
|
||||
self, mock_ensure, mock_registered, mock_context,
|
||||
):
|
||||
"""Verify multiple stages are initialized in order."""
|
||||
stage1_cls = MagicMock()
|
||||
stage1_cls.__name__ = "Stage1"
|
||||
stage1_instance = MagicMock()
|
||||
stage1_instance.initialize = AsyncMock()
|
||||
stage1_cls.return_value = stage1_instance
|
||||
|
||||
stage2_cls = MagicMock()
|
||||
stage2_cls.__name__ = "Stage2"
|
||||
stage2_instance = MagicMock()
|
||||
stage2_instance.initialize = AsyncMock()
|
||||
stage2_cls.return_value = stage2_instance
|
||||
|
||||
mock_registered.extend([stage1_cls, stage2_cls])
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
|
||||
await scheduler.initialize()
|
||||
|
||||
assert len(scheduler.stages) == 2
|
||||
stage1_instance.initialize.assert_awaited_once_with(mock_context)
|
||||
stage2_instance.initialize.assert_awaited_once_with(mock_context)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _process_stages
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPipelineSchedulerProcessStages:
|
||||
"""Tests for PipelineScheduler._process_stages()."""
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_generator_stage_executed(
|
||||
self, mock_ensure, mock_registered, mock_context, mock_stage_cls,
|
||||
):
|
||||
"""Verify a non-generator stage is awaited."""
|
||||
mock_registered.append(mock_stage_cls)
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = [mock_stage_cls.return_value]
|
||||
event = MagicMock()
|
||||
event.is_stopped.return_value = False
|
||||
|
||||
await scheduler._process_stages(event)
|
||||
|
||||
mock_stage_cls.return_value.process.assert_called_once_with(event)
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_generator_stage_yields_and_next_stage_runs(
|
||||
self, mock_ensure, mock_registered, mock_context,
|
||||
):
|
||||
"""Verify a generator stage yields and the next stage runs."""
|
||||
async def gen_process(_event):
|
||||
yield None
|
||||
|
||||
stage1 = MagicMock()
|
||||
stage1.process = gen_process
|
||||
stage1.__class__.__name__ = "Stage1"
|
||||
|
||||
stage2 = MagicMock()
|
||||
stage2.process = AsyncMock()
|
||||
stage2.__class__.__name__ = "Stage2"
|
||||
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = [stage1, stage2]
|
||||
event = MagicMock()
|
||||
event.is_stopped.return_value = False
|
||||
|
||||
await scheduler._process_stages(event)
|
||||
|
||||
stage2.process.assert_called_once_with(event)
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_generator_stage_stops_propagation(
|
||||
self, mock_ensure, mock_registered, mock_context,
|
||||
):
|
||||
"""Verify event.stop_event() breaks the pipeline in generator stage."""
|
||||
stage1_pass = [False]
|
||||
|
||||
async def onion_process(_event):
|
||||
stage1_pass[0] = True
|
||||
yield None
|
||||
|
||||
stage1 = MagicMock()
|
||||
stage1.process = onion_process
|
||||
stage1.__class__.__name__ = "Stage1"
|
||||
|
||||
stage2 = MagicMock()
|
||||
stage2.process = AsyncMock()
|
||||
stage2.__class__.__name__ = "Stage2"
|
||||
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = [stage1, stage2]
|
||||
event = MagicMock()
|
||||
event.is_stopped = lambda: True # always stopped
|
||||
|
||||
await scheduler._process_stages(event)
|
||||
|
||||
stage2.process.assert_not_called()
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_generator_stage_stops_propagation(
|
||||
self, mock_ensure, mock_registered, mock_context,
|
||||
):
|
||||
"""Verify event.stop_event() breaks non-generator stage chain."""
|
||||
stage1 = MagicMock()
|
||||
stage1.process = AsyncMock()
|
||||
stage1.__class__.__name__ = "Stage1"
|
||||
|
||||
stage2 = MagicMock()
|
||||
stage2.process = AsyncMock()
|
||||
stage2.__class__.__name__ = "Stage2"
|
||||
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = [stage1, stage2]
|
||||
event = MagicMock()
|
||||
event.is_stopped = lambda: True # always stopped
|
||||
|
||||
await scheduler._process_stages(event)
|
||||
|
||||
stage1.process.assert_called_once_with(event)
|
||||
stage2.process.assert_not_called()
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_generator_with_onion_recursion(
|
||||
self, mock_ensure, mock_registered, mock_context,
|
||||
):
|
||||
"""Verify generator stages recursively process subsequent stages (onion model)."""
|
||||
call_order = []
|
||||
|
||||
async def onion_process(event):
|
||||
call_order.append("before_yield")
|
||||
yield None
|
||||
call_order.append("after_yield")
|
||||
|
||||
stage1 = MagicMock()
|
||||
stage1.process = onion_process
|
||||
stage1.__class__.__name__ = "Stage1"
|
||||
|
||||
stage2 = MagicMock()
|
||||
stage2.process = AsyncMock(side_effect=lambda e: call_order.append("stage2"))
|
||||
stage2.__class__.__name__ = "Stage2"
|
||||
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = [stage1, stage2]
|
||||
event = MagicMock()
|
||||
event.is_stopped.return_value = False
|
||||
|
||||
await scheduler._process_stages(event)
|
||||
|
||||
# Order: stage1 before_yield -> yield -> stage2 -> stage1 after_yield
|
||||
assert call_order == ["before_yield", "stage2", "after_yield"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# execute
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPipelineSchedulerExecute:
|
||||
"""Tests for PipelineScheduler.execute()."""
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.active_event_registry")
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_calls_process_stages(
|
||||
self, mock_ensure, mock_registered, mock_registry, mock_context,
|
||||
):
|
||||
"""Verify execute calls _process_stages and cleans up."""
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = []
|
||||
scheduler._process_stages = AsyncMock()
|
||||
|
||||
event = MagicMock()
|
||||
await scheduler.execute(event)
|
||||
|
||||
scheduler._process_stages.assert_awaited_once_with(event)
|
||||
mock_registry.register.assert_called_once_with(event)
|
||||
mock_registry.unregister.assert_called_once_with(event)
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.active_event_registry")
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_webchat_event_sends_none(
|
||||
self, mock_ensure, mock_registered, mock_registry, mock_context,
|
||||
):
|
||||
"""Verify WebChatMessageEvent gets an extra None send."""
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = []
|
||||
scheduler._process_stages = AsyncMock()
|
||||
|
||||
event = MagicMock(spec=["send", "is_stopped"])
|
||||
event.__class__.__name__ = "WebChatMessageEvent"
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.scheduler.WebChatMessageEvent",
|
||||
event.__class__,
|
||||
):
|
||||
await scheduler.execute(event)
|
||||
|
||||
event.send.assert_awaited_once_with(None)
|
||||
mock_registry.unregister.assert_called_once_with(event)
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.active_event_registry")
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_wecom_event_sends_none(
|
||||
self, mock_ensure, mock_registered, mock_registry, mock_context,
|
||||
):
|
||||
"""Verify WecomAIBotMessageEvent gets an extra None send."""
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = []
|
||||
scheduler._process_stages = AsyncMock()
|
||||
|
||||
event = MagicMock(spec=["send", "is_stopped"])
|
||||
event.__class__.__name__ = "WecomAIBotMessageEvent"
|
||||
|
||||
with patch(
|
||||
"astrbot.core.pipeline.scheduler.WecomAIBotMessageEvent",
|
||||
event.__class__,
|
||||
):
|
||||
await scheduler.execute(event)
|
||||
|
||||
event.send.assert_awaited_once_with(None)
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.active_event_registry")
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_normal_event_no_extra_send(
|
||||
self, mock_ensure, mock_registered, mock_registry, mock_context,
|
||||
):
|
||||
"""Verify regular events do not get extra None send."""
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = []
|
||||
scheduler._process_stages = AsyncMock()
|
||||
|
||||
event = MagicMock()
|
||||
event.__class__.__name__ = "NormalEvent"
|
||||
|
||||
await scheduler.execute(event)
|
||||
|
||||
event.send.assert_not_called()
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.active_event_registry")
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_with_sdk_plugin_bridge(
|
||||
self, mock_ensure, mock_registered, mock_registry, mock_context,
|
||||
):
|
||||
"""Verify sdk_plugin_bridge.close_request_overlay_for_event is called."""
|
||||
mock_bridge = MagicMock()
|
||||
mock_bridge.close_request_overlay_for_event = MagicMock()
|
||||
mock_context.plugin_manager.context.sdk_plugin_bridge = mock_bridge
|
||||
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = []
|
||||
scheduler._process_stages = AsyncMock()
|
||||
|
||||
event = MagicMock()
|
||||
await scheduler.execute(event)
|
||||
|
||||
mock_bridge.close_request_overlay_for_event.assert_called_once_with(event)
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.active_event_registry")
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_without_sdk_plugin_bridge(
|
||||
self, mock_ensure, mock_registered, mock_registry, mock_context,
|
||||
):
|
||||
"""Verify no error when sdk_plugin_bridge is absent."""
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = []
|
||||
scheduler._process_stages = AsyncMock()
|
||||
|
||||
event = MagicMock()
|
||||
await scheduler.execute(event)
|
||||
|
||||
# Should not raise
|
||||
|
||||
@patch("astrbot.core.pipeline.scheduler.active_event_registry")
|
||||
@patch("astrbot.core.pipeline.scheduler.registered_stages", [])
|
||||
@patch("astrbot.core.pipeline.scheduler.ensure_builtin_stages_registered")
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_unregisters_on_error(
|
||||
self, mock_ensure, mock_registered, mock_registry, mock_context,
|
||||
):
|
||||
"""Verify event is still unregistered when _process_stages raises."""
|
||||
scheduler = PipelineScheduler(mock_context)
|
||||
scheduler.stages = []
|
||||
scheduler._process_stages = AsyncMock(side_effect=RuntimeError("fail"))
|
||||
|
||||
event = MagicMock()
|
||||
with pytest.raises(RuntimeError):
|
||||
await scheduler.execute(event)
|
||||
|
||||
mock_registry.unregister.assert_called_once_with(event)
|
||||
@@ -0,0 +1,312 @@
|
||||
"""Tests for astrbot.core.platform.platform — Platform ABC."""
|
||||
|
||||
import asyncio
|
||||
from asyncio import Queue
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.platform.platform import Platform, PlatformError, PlatformStatus
|
||||
from astrbot.core.platform.platform_metadata import PlatformMetadata
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Concrete subclass so we can test non-abstract behaviour of Platform.
|
||||
# ---------------------------------------------------------------------------
|
||||
class ConcretePlatform(Platform):
|
||||
"""Minimal concrete Platform used in every test that needs an instance."""
|
||||
|
||||
async def run(self) -> None:
|
||||
pass
|
||||
|
||||
def meta(self) -> PlatformMetadata:
|
||||
return PlatformMetadata(
|
||||
name="test_adapter",
|
||||
description="A test adapter",
|
||||
id="test_adapter_id",
|
||||
adapter_display_name="Test Adapter",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.fixture
|
||||
def config() -> dict:
|
||||
return {
|
||||
"key": "value",
|
||||
"unified_webhook_mode": False,
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def event_queue() -> Queue:
|
||||
return Queue()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def platform(config: dict, event_queue: Queue) -> ConcretePlatform:
|
||||
return ConcretePlatform(config, event_queue)
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Construction
|
||||
# ===================================================================
|
||||
|
||||
class TestConstruction:
|
||||
"""Platform.__init__ stores constructor arguments and sets defaults."""
|
||||
|
||||
def test_stores_config_and_queue(self, config: dict, event_queue: Queue):
|
||||
p = ConcretePlatform(config, event_queue)
|
||||
assert p.config is config
|
||||
assert p._event_queue is event_queue
|
||||
|
||||
def test_default_status_is_pending(self, platform: ConcretePlatform):
|
||||
assert platform.status == PlatformStatus.PENDING
|
||||
|
||||
def test_client_self_id_is_random_hex(self, platform: ConcretePlatform):
|
||||
assert isinstance(platform.client_self_id, str)
|
||||
assert len(platform.client_self_id) == 32 # uuid4 hex
|
||||
|
||||
def test_errors_list_starts_empty(self, platform: ConcretePlatform):
|
||||
assert platform.errors == []
|
||||
|
||||
def test_started_at_is_none_until_running(self, platform: ConcretePlatform):
|
||||
assert platform._started_at is None
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Status property
|
||||
# ===================================================================
|
||||
|
||||
class TestStatus:
|
||||
"""Platform.status getter/setter and side-effects."""
|
||||
|
||||
def test_setter_changes_status(self, platform: ConcretePlatform):
|
||||
platform.status = PlatformStatus.RUNNING
|
||||
assert platform.status == PlatformStatus.RUNNING
|
||||
|
||||
def test_setting_running_records_started_at(self, platform: ConcretePlatform):
|
||||
platform.status = PlatformStatus.RUNNING
|
||||
assert platform._started_at is not None
|
||||
assert isinstance(platform._started_at, datetime)
|
||||
|
||||
def test_setting_running_twice_does_not_overwrite_started_at(
|
||||
self, platform: ConcretePlatform
|
||||
):
|
||||
platform.status = PlatformStatus.RUNNING
|
||||
t1 = platform._started_at
|
||||
platform.status = PlatformStatus.ERROR
|
||||
platform.status = PlatformStatus.RUNNING
|
||||
assert platform._started_at == t1
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Error paths
|
||||
# ===================================================================
|
||||
|
||||
class TestErrors:
|
||||
"""record_error, last_error, clear_errors."""
|
||||
|
||||
def test_record_error_appends_to_list(self, platform: ConcretePlatform):
|
||||
platform.record_error("something went wrong", "traceback line 1")
|
||||
assert len(platform.errors) == 1
|
||||
assert platform.errors[0].message == "something went wrong"
|
||||
assert platform.errors[0].traceback == "traceback line 1"
|
||||
|
||||
def test_record_error_sets_status_to_error(self, platform: ConcretePlatform):
|
||||
platform.record_error("fail")
|
||||
assert platform.status == PlatformStatus.ERROR
|
||||
|
||||
def test_last_error_none_when_empty(self, platform: ConcretePlatform):
|
||||
assert platform.last_error is None
|
||||
|
||||
def test_last_error_returns_most_recent(self, platform: ConcretePlatform):
|
||||
platform.record_error("first")
|
||||
platform.record_error("second")
|
||||
assert platform.last_error.message == "second"
|
||||
|
||||
def test_clear_errors_empties_list(self, platform: ConcretePlatform):
|
||||
platform.record_error("first")
|
||||
platform.clear_errors()
|
||||
assert platform.errors == []
|
||||
|
||||
def test_clear_errors_resets_status_from_error(self, platform: ConcretePlatform):
|
||||
platform.record_error("first")
|
||||
platform.clear_errors()
|
||||
assert platform.status == PlatformStatus.RUNNING
|
||||
|
||||
def test_clear_errors_is_noop_when_no_error(self, platform: ConcretePlatform):
|
||||
platform.clear_errors()
|
||||
assert platform.errors == []
|
||||
assert platform.status == PlatformStatus.PENDING
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# unified_webhook
|
||||
# ===================================================================
|
||||
|
||||
class TestUnifiedWebhook:
|
||||
"""Platform.unified_webhook() logic."""
|
||||
|
||||
def test_disabled_by_default(self, platform: ConcretePlatform):
|
||||
assert platform.unified_webhook() is False
|
||||
|
||||
def test_enabled_when_both_config_present(
|
||||
self, config: dict, event_queue: Queue
|
||||
):
|
||||
config["unified_webhook_mode"] = True
|
||||
config["webhook_uuid"] = "abc123"
|
||||
p = ConcretePlatform(config, event_queue)
|
||||
assert p.unified_webhook() is True
|
||||
|
||||
def test_disabled_without_uuid(self, config: dict, event_queue: Queue):
|
||||
config["unified_webhook_mode"] = True
|
||||
# no webhook_uuid set
|
||||
p = ConcretePlatform(config, event_queue)
|
||||
assert p.unified_webhook() is False
|
||||
|
||||
def test_disabled_when_mode_off(self, config: dict, event_queue: Queue):
|
||||
config["unified_webhook_mode"] = False
|
||||
config["webhook_uuid"] = "abc123"
|
||||
p = ConcretePlatform(config, event_queue)
|
||||
assert p.unified_webhook() is False
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# get_stats
|
||||
# ===================================================================
|
||||
|
||||
class TestGetStats:
|
||||
"""Platform.get_stats() structure and content."""
|
||||
|
||||
def test_stats_contains_expected_keys(self, platform: ConcretePlatform):
|
||||
stats = platform.get_stats()
|
||||
expected_keys = {
|
||||
"id", "type", "display_name", "status", "started_at",
|
||||
"error_count", "last_error", "unified_webhook", "meta",
|
||||
}
|
||||
assert expected_keys.issubset(stats.keys())
|
||||
|
||||
def test_stats_values_without_errors(self, platform: ConcretePlatform):
|
||||
stats = platform.get_stats()
|
||||
assert stats["id"] == "test_adapter_id"
|
||||
assert stats["type"] == "test_adapter"
|
||||
assert stats["display_name"] == "Test Adapter"
|
||||
assert stats["status"] == PlatformStatus.PENDING.value
|
||||
assert stats["started_at"] is None
|
||||
assert stats["error_count"] == 0
|
||||
assert stats["last_error"] is None
|
||||
assert stats["unified_webhook"] is False
|
||||
assert stats["meta"]["id"] == "test_adapter_id"
|
||||
assert stats["meta"]["name"] == "test_adapter"
|
||||
|
||||
def test_stats_reflects_recorded_errors(self, platform: ConcretePlatform):
|
||||
platform.record_error("err1", "tb1")
|
||||
platform.record_error("err2", "tb2")
|
||||
stats = platform.get_stats()
|
||||
assert stats["error_count"] == 2
|
||||
assert stats["last_error"]["message"] == "err2"
|
||||
assert stats["last_error"]["traceback"] == "tb2"
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Instance methods
|
||||
# ===================================================================
|
||||
|
||||
class TestMethods:
|
||||
"""terminate, get_client, commit_event, webhook_callback, send_by_session."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminate_sets_stopped(self, platform: ConcretePlatform):
|
||||
await platform.terminate()
|
||||
assert platform.status == PlatformStatus.STOPPED
|
||||
|
||||
def test_get_client_returns_none(self, platform: ConcretePlatform):
|
||||
assert platform.get_client() is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_event_puts_into_queue(
|
||||
self, platform: ConcretePlatform, event_queue: Queue
|
||||
):
|
||||
mock_event = MagicMock()
|
||||
platform.commit_event(mock_event)
|
||||
assert event_queue.qsize() == 1
|
||||
assert await event_queue.get() is mock_event
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_webhook_callback_raises_not_implemented(
|
||||
self, platform: ConcretePlatform
|
||||
):
|
||||
with pytest.raises(NotImplementedError) as exc:
|
||||
await platform.webhook_callback(None)
|
||||
assert "未实现统一 Webhook 模式" in str(exc.value)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_by_session_calls_metric_upload(
|
||||
self, platform: ConcretePlatform
|
||||
):
|
||||
mock_session = MagicMock()
|
||||
mock_chain = MagicMock()
|
||||
with patch(
|
||||
"astrbot.core.platform.platform.Metric.upload",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_upload:
|
||||
await platform.send_by_session(mock_session, mock_chain)
|
||||
mock_upload.assert_called_once_with(
|
||||
msg_event_tick=1, adapter_name="test_adapter"
|
||||
)
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# Abstract-method detection (cannot instantiate ABC directly)
|
||||
# ===================================================================
|
||||
|
||||
class TestAbstractDetection:
|
||||
"""Verify the ABC enforces that run() and meta() are implemented."""
|
||||
|
||||
def test_cannot_instantiate_platform_directly(self):
|
||||
with pytest.raises(TypeError):
|
||||
Platform({"k": "v"}, Queue()) # type: ignore[abstract]
|
||||
|
||||
def test_missing_run_raises_type_error(self):
|
||||
class MissingRun(Platform):
|
||||
def meta(self) -> PlatformMetadata:
|
||||
return PlatformMetadata(name="x", description="x", id="x")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
MissingRun({}, Queue()) # type: ignore[abstract]
|
||||
|
||||
def test_missing_meta_raises_type_error(self):
|
||||
class MissingMeta(Platform):
|
||||
async def run(self) -> None:
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
MissingMeta({}, Queue()) # type: ignore[abstract]
|
||||
|
||||
|
||||
# ===================================================================
|
||||
# PlatformError dataclass
|
||||
# ===================================================================
|
||||
|
||||
class TestPlatformErrorDataclass:
|
||||
"""PlatformError construction and defaults."""
|
||||
|
||||
def test_required_message(self):
|
||||
err = PlatformError(message="oops")
|
||||
assert err.message == "oops"
|
||||
assert err.traceback is None
|
||||
assert isinstance(err.timestamp, datetime)
|
||||
|
||||
def test_with_traceback(self):
|
||||
err = PlatformError(message="oops", traceback="tb content")
|
||||
assert err.traceback == "tb content"
|
||||
|
||||
def test_default_timestamp_is_nowish(self):
|
||||
before = datetime.now()
|
||||
err = PlatformError(message="oops")
|
||||
after = datetime.now()
|
||||
assert before <= err.timestamp <= after
|
||||
@@ -0,0 +1,518 @@
|
||||
"""Tests for astrbot.core.provider.entities.
|
||||
|
||||
Covers ProviderRequest, LLMResponse, TokenUsage, ToolCallsResult,
|
||||
ProviderMeta, ProviderMetaData, and RerankResult construction and edge cases.
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.agent.message import (
|
||||
AssistantMessageSegment,
|
||||
ContentPart,
|
||||
ToolCall,
|
||||
ToolCallMessageSegment,
|
||||
)
|
||||
from astrbot.core.message.message_event_result import MessageChain
|
||||
from astrbot.core.provider.entities import (
|
||||
LLMResponse,
|
||||
ProviderMeta,
|
||||
ProviderMetaData,
|
||||
ProviderRequest,
|
||||
ProviderType,
|
||||
RerankResult,
|
||||
TokenUsage,
|
||||
ToolCallsResult,
|
||||
)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# ProviderMeta / ProviderMetaData
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestProviderMeta:
|
||||
def test_basic_construction(self):
|
||||
meta = ProviderMeta(id="p1", model="gpt-4", type="openai")
|
||||
assert meta.id == "p1"
|
||||
assert meta.model == "gpt-4"
|
||||
assert meta.type == "openai"
|
||||
assert meta.provider_type == ProviderType.CHAT_COMPLETION
|
||||
|
||||
def test_construction_with_provider_type(self):
|
||||
meta = ProviderMeta(
|
||||
id="emb1",
|
||||
model="text-embedding-3",
|
||||
type="openai_embedding",
|
||||
provider_type=ProviderType.EMBEDDING,
|
||||
)
|
||||
assert meta.provider_type == ProviderType.EMBEDDING
|
||||
|
||||
|
||||
class TestProviderMetaData:
|
||||
def test_basic_construction(self):
|
||||
pmd = ProviderMetaData(
|
||||
id="p1", model=None, type="openai", desc="OpenAI provider"
|
||||
)
|
||||
assert pmd.id == "p1"
|
||||
assert pmd.desc == "OpenAI provider"
|
||||
assert pmd.cls_type is None
|
||||
assert pmd.default_config_tmpl is None
|
||||
assert pmd.provider_display_name is None
|
||||
|
||||
def test_construction_with_all_fields(self):
|
||||
fake_cls = type("FakeProvider", (), {})
|
||||
pmd = ProviderMetaData(
|
||||
id="p2",
|
||||
model="gpt-4o",
|
||||
type="openai",
|
||||
desc="desc",
|
||||
provider_type=ProviderType.CHAT_COMPLETION,
|
||||
cls_type=fake_cls,
|
||||
default_config_tmpl={"key": "val"},
|
||||
provider_display_name="OpenAI Official",
|
||||
)
|
||||
assert pmd.cls_type is fake_cls
|
||||
assert pmd.default_config_tmpl == {"key": "val"}
|
||||
assert pmd.provider_display_name == "OpenAI Official"
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# TokenUsage
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestTokenUsage:
|
||||
def test_defaults(self):
|
||||
tu = TokenUsage()
|
||||
assert tu.input_other == 0
|
||||
assert tu.input_cached == 0
|
||||
assert tu.output == 0
|
||||
assert tu.total == 0
|
||||
assert tu.input == 0
|
||||
|
||||
def test_properties(self):
|
||||
tu = TokenUsage(input_other=10, input_cached=5, output=20)
|
||||
assert tu.total == 35
|
||||
assert tu.input == 15
|
||||
|
||||
def test_addition(self):
|
||||
a = TokenUsage(input_other=5, input_cached=2, output=10)
|
||||
b = TokenUsage(input_other=3, input_cached=1, output=4)
|
||||
result = a + b
|
||||
assert result.input_other == 8
|
||||
assert result.input_cached == 3
|
||||
assert result.output == 14
|
||||
|
||||
def test_subtraction(self):
|
||||
a = TokenUsage(input_other=10, input_cached=5, output=20)
|
||||
b = TokenUsage(input_other=3, input_cached=2, output=5)
|
||||
result = a - b
|
||||
assert result.input_other == 7
|
||||
assert result.input_cached == 3
|
||||
assert result.output == 15
|
||||
|
||||
def test_addition_preserves_immutability(self):
|
||||
a = TokenUsage(input_other=1, output=2)
|
||||
b = TokenUsage(input_other=3, output=4)
|
||||
c = a + b
|
||||
assert a.input_other == 1
|
||||
assert b.input_other == 3
|
||||
assert c.input_other == 4
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# ToolCallsResult
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestToolCallsResult:
|
||||
def test_construction(self):
|
||||
info = MagicMock(spec=AssistantMessageSegment)
|
||||
info.model_dump.return_value = {"role": "assistant", "content": "thinking"}
|
||||
result_seg = MagicMock(spec=ToolCallMessageSegment)
|
||||
result_seg.model_dump.return_value = {"role": "tool", "content": "result"}
|
||||
|
||||
tcr = ToolCallsResult(
|
||||
tool_calls_info=info,
|
||||
tool_calls_result=[result_seg],
|
||||
)
|
||||
assert tcr.tool_calls_info is info
|
||||
assert len(tcr.tool_calls_result) == 1
|
||||
|
||||
def test_to_openai_messages(self):
|
||||
info = MagicMock(spec=AssistantMessageSegment)
|
||||
info.model_dump.return_value = {"role": "assistant"}
|
||||
r1 = MagicMock(spec=ToolCallMessageSegment)
|
||||
r1.model_dump.return_value = {"role": "tool", "name": "get_weather"}
|
||||
r2 = MagicMock(spec=ToolCallMessageSegment)
|
||||
r2.model_dump.return_value = {"role": "tool", "name": "search"}
|
||||
|
||||
tcr = ToolCallsResult(tool_calls_info=info, tool_calls_result=[r1, r2])
|
||||
msgs = tcr.to_openai_messages()
|
||||
assert len(msgs) == 3
|
||||
assert msgs[0] == {"role": "assistant"}
|
||||
assert msgs[1] == {"role": "tool", "name": "get_weather"}
|
||||
assert msgs[2] == {"role": "tool", "name": "search"}
|
||||
|
||||
def test_to_openai_messages_model(self):
|
||||
info = MagicMock(spec=AssistantMessageSegment)
|
||||
r1 = MagicMock(spec=ToolCallMessageSegment)
|
||||
|
||||
tcr = ToolCallsResult(tool_calls_info=info, tool_calls_result=[r1])
|
||||
models = tcr.to_openai_messages_model()
|
||||
assert len(models) == 2
|
||||
assert models[0] is info
|
||||
assert models[1] is r1
|
||||
|
||||
def test_to_openai_messages_empty_result(self):
|
||||
info = MagicMock(spec=AssistantMessageSegment)
|
||||
info.model_dump.return_value = {"role": "assistant"}
|
||||
tcr = ToolCallsResult(tool_calls_info=info, tool_calls_result=[])
|
||||
msgs = tcr.to_openai_messages()
|
||||
assert len(msgs) == 1
|
||||
assert msgs[0] == {"role": "assistant"}
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# ProviderRequest
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestProviderRequest:
|
||||
def test_default_construction(self):
|
||||
req = ProviderRequest()
|
||||
assert req.prompt is None
|
||||
assert req.session_id == ""
|
||||
assert req.image_urls == []
|
||||
assert req.audio_urls == []
|
||||
assert req.contexts == []
|
||||
assert req.func_tool is None
|
||||
assert req.system_prompt is None
|
||||
assert req.conversation is None
|
||||
assert req.tool_calls_result is None
|
||||
assert req.model is None
|
||||
|
||||
def test_construction_with_values(self):
|
||||
req = ProviderRequest(
|
||||
prompt="Hello",
|
||||
session_id="sess-1",
|
||||
image_urls=["http://example.com/img.png"],
|
||||
system_prompt="You are a bot",
|
||||
model="gpt-4",
|
||||
)
|
||||
assert req.prompt == "Hello"
|
||||
assert req.session_id == "sess-1"
|
||||
assert req.image_urls == ["http://example.com/img.png"]
|
||||
assert req.system_prompt == "You are a bot"
|
||||
assert req.model == "gpt-4"
|
||||
|
||||
def test_repr_without_context(self):
|
||||
req = ProviderRequest(prompt="hi", session_id="s1")
|
||||
text = repr(req)
|
||||
assert "prompt=hi" in text
|
||||
assert "session_id=s1" in text
|
||||
assert "image_count=0" in text
|
||||
|
||||
def test_repr_with_conversation(self):
|
||||
conversation = MagicMock()
|
||||
conversation.cid = "conv-abc"
|
||||
req = ProviderRequest(prompt="test", conversation=conversation)
|
||||
text = repr(req)
|
||||
assert "conversation_id=conv-abc" in text
|
||||
|
||||
def test_append_tool_calls_result_none_to_single(self):
|
||||
req = ProviderRequest()
|
||||
tcr = MagicMock(spec=ToolCallsResult)
|
||||
req.append_tool_calls_result(tcr)
|
||||
assert isinstance(req.tool_calls_result, list)
|
||||
assert len(req.tool_calls_result) == 1
|
||||
assert req.tool_calls_result[0] is tcr
|
||||
|
||||
def test_append_tool_calls_result_single_to_list(self):
|
||||
tcr1 = MagicMock(spec=ToolCallsResult)
|
||||
req = ProviderRequest(tool_calls_result=tcr1)
|
||||
tcr2 = MagicMock(spec=ToolCallsResult)
|
||||
req.append_tool_calls_result(tcr2)
|
||||
assert isinstance(req.tool_calls_result, list)
|
||||
assert len(req.tool_calls_result) == 2
|
||||
assert req.tool_calls_result[0] is tcr1
|
||||
assert req.tool_calls_result[1] is tcr2
|
||||
|
||||
def test_append_tool_calls_result_list(self):
|
||||
tcr1 = MagicMock(spec=ToolCallsResult)
|
||||
tcr2 = MagicMock(spec=ToolCallsResult)
|
||||
req = ProviderRequest(tool_calls_result=[tcr1])
|
||||
tcr3 = MagicMock(spec=ToolCallsResult)
|
||||
req.append_tool_calls_result(tcr3)
|
||||
assert len(req.tool_calls_result) == 2
|
||||
|
||||
def test_print_friendly_context_no_contexts(self):
|
||||
req = ProviderRequest(prompt="hello", image_urls=["a.png"], audio_urls=["b.wav"])
|
||||
result = req._print_friendly_context()
|
||||
assert "prompt: hello" in result
|
||||
assert "image_count: 1" in result
|
||||
assert "audio_count: 1" in result
|
||||
|
||||
def test_print_friendly_context_with_text_contexts(self):
|
||||
req = ProviderRequest(contexts=[{"role": "user", "content": "hello"}])
|
||||
result = req._print_friendly_context()
|
||||
assert "user: hello" in result
|
||||
|
||||
def test_print_friendly_context_filters_checkpoints(self):
|
||||
req = ProviderRequest(
|
||||
contexts=[
|
||||
{"role": "user", "content": "hi", "checkpoint": True},
|
||||
{"role": "assistant", "content": "hello"},
|
||||
]
|
||||
)
|
||||
with patch(
|
||||
"astrbot.core.provider.entities.is_checkpoint_message",
|
||||
side_effect=lambda c: c.get("checkpoint", False),
|
||||
):
|
||||
result = req._print_friendly_context()
|
||||
assert "user: hi" not in result
|
||||
assert "assistant: hello" in result
|
||||
|
||||
def test_print_friendly_context_multimodal(self):
|
||||
req = ProviderRequest(
|
||||
contexts=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "describe this"},
|
||||
{"type": "image_url", "image_url": {"url": "x.jpg"}},
|
||||
{"type": "image_url", "image_url": {"url": "y.jpg"}},
|
||||
{"type": "audio_url", "audio_url": {"url": "z.wav"}},
|
||||
],
|
||||
}
|
||||
]
|
||||
)
|
||||
result = req._print_friendly_context()
|
||||
assert "user: describe this[+2 images][+1 audios]" in result
|
||||
|
||||
def test_assemble_context_simple_text(self):
|
||||
"""When there's only a plain text prompt and no extra content, it returns a simple str content."""
|
||||
req = ProviderRequest(prompt="hello world")
|
||||
import asyncio
|
||||
|
||||
ctx = asyncio.run(req.assemble_context())
|
||||
assert ctx == {"role": "user", "content": "hello world"}
|
||||
|
||||
def test_assemble_context_empty_prompt_with_images_adds_placeholder(self):
|
||||
req = ProviderRequest(prompt="", image_urls=[])
|
||||
import asyncio
|
||||
|
||||
with patch.object(req, "_encode_image_bs64", return_value="data:image/jpeg;base64,abc"):
|
||||
req.image_urls = ["http://example.com/img.png"]
|
||||
ctx = asyncio.run(req.assemble_context())
|
||||
assert ctx["role"] == "user"
|
||||
# Should include "[图片]" placeholder
|
||||
content = ctx["content"]
|
||||
assert isinstance(content, list)
|
||||
assert any(b.get("text") == "[图片]" for b in content)
|
||||
|
||||
def test_assemble_context_empty_prompt_with_audio_adds_placeholder(self):
|
||||
req = ProviderRequest(prompt="", audio_urls=["http://example.com/a.wav"])
|
||||
import asyncio
|
||||
|
||||
with (
|
||||
patch.object(req, "_encode_audio_bs64", return_value="data:audio/wav;base64,xyz"),
|
||||
patch("astrbot.core.provider.entities.download_file", AsyncMock()),
|
||||
):
|
||||
ctx = asyncio.run(req.assemble_context())
|
||||
assert ctx["role"] == "user"
|
||||
content = ctx["content"]
|
||||
assert isinstance(content, list)
|
||||
assert any(b.get("text") == "[音频]" for b in content)
|
||||
|
||||
def test_assemble_context_with_extra_user_content(self):
|
||||
extra_part = MagicMock(spec=ContentPart)
|
||||
extra_part.model_dump.return_value = {"type": "text", "text": "extra instruction"}
|
||||
req = ProviderRequest(
|
||||
prompt="translate this",
|
||||
extra_user_content_parts=[extra_part],
|
||||
)
|
||||
import asyncio
|
||||
|
||||
ctx = asyncio.run(req.assemble_context())
|
||||
assert ctx["role"] == "user"
|
||||
content = ctx["content"]
|
||||
assert isinstance(content, list)
|
||||
texts = [b["text"] for b in content if b.get("type") == "text"]
|
||||
assert "translate this" in texts
|
||||
assert "extra instruction" in texts
|
||||
|
||||
def test_encode_image_bs64_base64_prefix(self):
|
||||
req = ProviderRequest()
|
||||
import asyncio
|
||||
|
||||
result = asyncio.run(req._encode_image_bs64("base64://rawdata"))
|
||||
assert result == "data:image/jpeg;base64,rawdata"
|
||||
|
||||
def test_encode_audio_bs64_base64_prefix(self):
|
||||
req = ProviderRequest()
|
||||
import asyncio
|
||||
|
||||
result = asyncio.run(req._encode_audio_bs64("base64://rawdata"))
|
||||
assert result == "data:audio/wav;base64,rawdata"
|
||||
|
||||
def test_str_equals_repr(self):
|
||||
req = ProviderRequest(prompt="test", session_id="s1")
|
||||
assert str(req) == repr(req)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# LLMResponse
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestLLMResponse:
|
||||
def test_default_role_assistant(self):
|
||||
resp = LLMResponse(role="assistant")
|
||||
assert resp.role == "assistant"
|
||||
assert resp.completion_text is None or resp.completion_text == ""
|
||||
assert resp.tools_call_args == []
|
||||
assert resp.tools_call_name == []
|
||||
assert resp.tools_call_ids == []
|
||||
assert resp.tools_call_extra_content == {}
|
||||
assert resp.reasoning_content is None
|
||||
assert resp.raw_completion is None
|
||||
assert resp.is_chunk is False
|
||||
assert resp.id is None
|
||||
assert resp.usage is None
|
||||
|
||||
def test_construction_with_completion_text(self):
|
||||
resp = LLMResponse(role="assistant", completion_text="Hello world")
|
||||
assert resp.completion_text == "Hello world"
|
||||
|
||||
def test_construction_with_result_chain(self):
|
||||
chain = MessageChain()
|
||||
chain.message("Hello from chain")
|
||||
resp = LLMResponse(role="assistant", result_chain=chain)
|
||||
assert resp.completion_text == "Hello from chain"
|
||||
|
||||
def test_completion_text_setter_with_result_chain(self):
|
||||
chain = MessageChain()
|
||||
chain.message("Old text")
|
||||
resp = LLMResponse(role="assistant", result_chain=chain)
|
||||
assert resp.completion_text == "Old text"
|
||||
resp.completion_text = "New text"
|
||||
# The setter inserts a Plain component at the start after removing old ones
|
||||
assert "New text" in resp.completion_text
|
||||
|
||||
def test_completion_text_setter_without_result_chain(self):
|
||||
resp = LLMResponse(role="assistant")
|
||||
resp.completion_text = "direct text"
|
||||
assert resp.completion_text == "direct text"
|
||||
|
||||
def test_tool_calls_defaults_to_empty(self):
|
||||
resp = LLMResponse(role="assistant")
|
||||
# They should be empty lists, not None
|
||||
assert resp.tools_call_args == []
|
||||
assert resp.tools_call_name == []
|
||||
assert resp.tools_call_ids == []
|
||||
assert resp.tools_call_extra_content == {}
|
||||
|
||||
def test_construction_with_tool_calls(self):
|
||||
resp = LLMResponse(
|
||||
role="assistant",
|
||||
tools_call_args=[{"location": "NYC"}],
|
||||
tools_call_name=["get_weather"],
|
||||
tools_call_ids=["call_123"],
|
||||
tools_call_extra_content={"call_123": {"source": "web"}},
|
||||
)
|
||||
assert resp.tools_call_args == [{"location": "NYC"}]
|
||||
assert resp.tools_call_name == ["get_weather"]
|
||||
assert resp.tools_call_ids == ["call_123"]
|
||||
assert resp.tools_call_extra_content == {"call_123": {"source": "web"}}
|
||||
|
||||
def test_to_openai_tool_calls(self):
|
||||
resp = LLMResponse(
|
||||
role="assistant",
|
||||
tools_call_args=[{"q": "weather"}, {"q": "news"}],
|
||||
tools_call_name=["search", "search"],
|
||||
tools_call_ids=["c1", "c2"],
|
||||
tools_call_extra_content={"c1": {"priority": 1}},
|
||||
)
|
||||
calls = resp.to_openai_tool_calls()
|
||||
assert len(calls) == 2
|
||||
assert calls[0]["id"] == "c1"
|
||||
assert calls[0]["function"]["name"] == "search"
|
||||
assert json.loads(calls[0]["function"]["arguments"]) == {"q": "weather"}
|
||||
assert calls[0]["extra_content"] == {"priority": 1}
|
||||
assert calls[1]["id"] == "c2"
|
||||
assert "extra_content" not in calls[1]
|
||||
|
||||
def test_to_openai_to_calls_model(self):
|
||||
resp = LLMResponse(
|
||||
role="assistant",
|
||||
tools_call_args=[{"x": 1}],
|
||||
tools_call_name=["foo"],
|
||||
tools_call_ids=["cid1"],
|
||||
tools_call_extra_content={"cid1": {"meta": "data"}},
|
||||
)
|
||||
calls = resp.to_openai_to_calls_model()
|
||||
assert len(calls) == 1
|
||||
assert isinstance(calls[0], ToolCall)
|
||||
assert calls[0].id == "cid1"
|
||||
assert calls[0].function.name == "foo"
|
||||
assert calls[0].extra_content == {"meta": "data"}
|
||||
|
||||
def test_to_openai_tool_calls_empty(self):
|
||||
resp = LLMResponse(role="assistant")
|
||||
calls = resp.to_openai_tool_calls()
|
||||
assert calls == []
|
||||
|
||||
def test_to_openai_to_calls_model_empty(self):
|
||||
resp = LLMResponse(role="assistant")
|
||||
calls = resp.to_openai_to_calls_model()
|
||||
assert calls == []
|
||||
|
||||
def test_construction_with_reasoning(self):
|
||||
resp = LLMResponse(
|
||||
role="assistant",
|
||||
completion_text="final answer",
|
||||
reasoning_content="thinking step by step",
|
||||
reasoning_signature="sig_abc",
|
||||
)
|
||||
assert resp.reasoning_content == "thinking step by step"
|
||||
assert resp.reasoning_signature == "sig_abc"
|
||||
|
||||
def test_construction_with_raw_completion(self):
|
||||
raw = MagicMock()
|
||||
resp = LLMResponse(role="assistant", raw_completion=raw)
|
||||
assert resp.raw_completion is raw
|
||||
|
||||
def test_construction_with_usage(self):
|
||||
usage = TokenUsage(input_other=10, output=20)
|
||||
resp = LLMResponse(role="assistant", usage=usage)
|
||||
assert resp.usage is usage
|
||||
|
||||
def test_construction_with_chunk_and_id(self):
|
||||
resp = LLMResponse(role="assistant", is_chunk=True, id="chunk_1")
|
||||
assert resp.is_chunk is True
|
||||
assert resp.id == "chunk_1"
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# RerankResult
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestRerankResult:
|
||||
def test_construction(self):
|
||||
rr = RerankResult(index=0, relevance_score=0.95)
|
||||
assert rr.index == 0
|
||||
assert rr.relevance_score == 0.95
|
||||
|
||||
def test_negative_score(self):
|
||||
rr = RerankResult(index=5, relevance_score=-0.1)
|
||||
assert rr.relevance_score == -0.1
|
||||
|
||||
def test_zero_values(self):
|
||||
rr = RerankResult(index=0, relevance_score=0.0)
|
||||
assert rr.index == 0
|
||||
assert rr.relevance_score == 0.0
|
||||
@@ -0,0 +1,408 @@
|
||||
"""Tests for astrbot.core.provider.manager.
|
||||
|
||||
Covers ProviderManager __init__, callbacks/hooks, get_provider_by_id,
|
||||
get_using_provider, _resolve_env_key_list, get_provider_config_by_id,
|
||||
and related helper methods.
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.provider.entities import ProviderType
|
||||
from astrbot.core.provider.manager import ProviderManager
|
||||
from astrbot.core.provider.provider import (
|
||||
EmbeddingProvider,
|
||||
Provider,
|
||||
RerankProvider,
|
||||
STTProvider,
|
||||
TTSProvider,
|
||||
)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Fixtures
|
||||
# =========================================================================
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_acm():
|
||||
"""Create a mock AstrBotConfigManager with valid nested config structure."""
|
||||
acm = MagicMock()
|
||||
acm.confs = {
|
||||
"default": {
|
||||
"provider": [],
|
||||
"provider_sources": [],
|
||||
"provider_settings": {"default_provider_id": ""},
|
||||
"provider_stt_settings": {},
|
||||
"provider_tts_settings": {},
|
||||
}
|
||||
}
|
||||
acm.default_conf = acm.confs["default"]
|
||||
acm.get_conf.return_value = acm.confs["default"]
|
||||
return acm
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_db():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_persona_mgr():
|
||||
pm = MagicMock()
|
||||
pm.default_persona = "default"
|
||||
pm.persona_v3_config = []
|
||||
pm.personas_v3 = []
|
||||
pm.selected_default_persona_v3 = None
|
||||
return pm
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def manager(mock_acm, mock_db, mock_persona_mgr):
|
||||
with patch("astrbot.core.provider.manager.llm_tools") as mock_llm_tools:
|
||||
mgr = ProviderManager(
|
||||
acm=mock_acm,
|
||||
db_helper=mock_db,
|
||||
persona_mgr=mock_persona_mgr,
|
||||
)
|
||||
yield mgr
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# __init__
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestProviderManagerInit:
|
||||
def test_construction(self, manager):
|
||||
assert manager.reload_lock is not None
|
||||
assert manager.resource_lock is not None
|
||||
assert manager.providers_config == []
|
||||
assert manager.provider_sources_config == []
|
||||
assert manager.provider_settings == {"default_provider_id": ""}
|
||||
assert manager.provider_insts == []
|
||||
assert manager.stt_provider_insts == []
|
||||
assert manager.tts_provider_insts == []
|
||||
assert manager.embedding_provider_insts == []
|
||||
assert manager.rerank_provider_insts == []
|
||||
assert manager.inst_map == {}
|
||||
assert manager.curr_provider_inst is None
|
||||
assert manager._provider_change_callback is None
|
||||
assert manager._provider_change_hooks == []
|
||||
assert manager._mcp_init_task is None
|
||||
|
||||
def test_default_persona_name_from_mgr(self, manager):
|
||||
assert manager.default_persona_name == "default"
|
||||
|
||||
def test_persona_configs_property(self, manager):
|
||||
assert manager.persona_configs == []
|
||||
|
||||
def test_personas_property(self, manager):
|
||||
assert manager.personas == []
|
||||
|
||||
def test_selected_default_persona_property(self, manager):
|
||||
assert manager.selected_default_persona is None
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Callbacks / Hooks
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestProviderManagerCallbacks:
|
||||
def test_set_provider_change_callback(self, manager):
|
||||
cb = MagicMock()
|
||||
manager.set_provider_change_callback(cb)
|
||||
assert manager._provider_change_callback is cb
|
||||
|
||||
def test_set_provider_change_callback_none(self, manager):
|
||||
manager.set_provider_change_callback(None)
|
||||
assert manager._provider_change_callback is None
|
||||
|
||||
def test_register_provider_change_hook(self, manager):
|
||||
hook = MagicMock()
|
||||
manager.register_provider_change_hook(hook)
|
||||
assert hook in manager._provider_change_hooks
|
||||
|
||||
def test_register_provider_change_hook_duplicate(self, manager):
|
||||
hook = MagicMock()
|
||||
manager.register_provider_change_hook(hook)
|
||||
manager.register_provider_change_hook(hook)
|
||||
assert len(manager._provider_change_hooks) == 1
|
||||
|
||||
def test_unregister_provider_change_hook(self, manager):
|
||||
hook = MagicMock()
|
||||
manager.register_provider_change_hook(hook)
|
||||
manager.unregister_provider_change_hook(hook)
|
||||
assert hook not in manager._provider_change_hooks
|
||||
|
||||
def test_unregister_provider_change_hook_not_registered(self, manager):
|
||||
hook = MagicMock()
|
||||
# Should not raise
|
||||
manager.unregister_provider_change_hook(hook)
|
||||
|
||||
def test_notify_provider_changed_calls_callback(self, manager):
|
||||
cb = MagicMock()
|
||||
manager.set_provider_change_callback(cb)
|
||||
manager._notify_provider_changed("p1", ProviderType.CHAT_COMPLETION, "umo_1")
|
||||
cb.assert_called_once_with("p1", ProviderType.CHAT_COMPLETION, "umo_1")
|
||||
|
||||
def test_notify_provider_changed_swallows_callback_error(self, manager):
|
||||
cb = MagicMock(side_effect=ValueError("oops"))
|
||||
manager.set_provider_change_callback(cb)
|
||||
# Should not raise
|
||||
manager._notify_provider_changed("p1", ProviderType.CHAT_COMPLETION, None)
|
||||
|
||||
def test_notify_provider_changed_calls_hooks(self, manager):
|
||||
hook1 = MagicMock()
|
||||
hook2 = MagicMock()
|
||||
manager.register_provider_change_hook(hook1)
|
||||
manager.register_provider_change_hook(hook2)
|
||||
manager._notify_provider_changed("p1", ProviderType.SPEECH_TO_TEXT, None)
|
||||
hook1.assert_called_once_with("p1", ProviderType.SPEECH_TO_TEXT, None)
|
||||
hook2.assert_called_once_with("p1", ProviderType.SPEECH_TO_TEXT, None)
|
||||
|
||||
def test_notify_provider_changed_skips_callback_in_hooks(self, manager):
|
||||
"""When the same callable is both callback and hook, it should only be invoked once."""
|
||||
fn = MagicMock()
|
||||
manager.set_provider_change_callback(fn)
|
||||
manager.register_provider_change_hook(fn)
|
||||
manager._notify_provider_changed("p1", ProviderType.CHAT_COMPLETION, None)
|
||||
assert fn.call_count == 1
|
||||
|
||||
def test_notify_provider_changed_swallows_hook_error(self, manager):
|
||||
hook = MagicMock(side_effect=RuntimeError("hook failed"))
|
||||
manager.register_provider_change_hook(hook)
|
||||
# Should not raise
|
||||
manager._notify_provider_changed("p1", ProviderType.CHAT_COMPLETION, None)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# get_provider_by_id / get_using_provider
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestProviderManagerLookups:
|
||||
def test_get_provider_by_id_not_found(self, manager):
|
||||
result = manager.get_provider_by_id("nonexistent")
|
||||
assert result is None
|
||||
|
||||
def test_get_provider_by_id_found(self, manager):
|
||||
fake_inst = MagicMock(spec=Provider)
|
||||
manager.inst_map["p1"] = fake_inst
|
||||
result = manager.get_provider_by_id("p1")
|
||||
assert result is fake_inst
|
||||
|
||||
def test_get_using_provider_chat_completion_default(self, manager, mock_acm):
|
||||
fake_provider = MagicMock(spec=Provider)
|
||||
manager.provider_insts = [fake_provider]
|
||||
mock_acm.get_conf.return_value = {
|
||||
"provider_settings": {"default_provider_id": None}
|
||||
}
|
||||
result = manager.get_using_provider(ProviderType.CHAT_COMPLETION)
|
||||
assert result is fake_provider
|
||||
|
||||
def test_get_using_provider_chat_completion_by_id(self, manager, mock_acm):
|
||||
fake_provider = MagicMock(spec=Provider)
|
||||
manager.inst_map["default_prov"] = fake_provider
|
||||
mock_acm.get_conf.return_value = {
|
||||
"provider_settings": {"default_provider_id": "default_prov"}
|
||||
}
|
||||
result = manager.get_using_provider(ProviderType.CHAT_COMPLETION)
|
||||
assert result is fake_provider
|
||||
|
||||
def test_get_using_provider_chat_completion_no_instances(self, manager, mock_acm):
|
||||
mock_acm.get_conf.return_value = {
|
||||
"provider_settings": {"default_provider_id": None}
|
||||
}
|
||||
result = manager.get_using_provider(ProviderType.CHAT_COMPLETION)
|
||||
assert result is None
|
||||
|
||||
def test_get_using_provider_stt_disabled_returns_none(self, manager, mock_acm):
|
||||
mock_acm.get_conf.return_value = {
|
||||
"provider_stt_settings": {"enable": False}
|
||||
}
|
||||
result = manager.get_using_provider(ProviderType.SPEECH_TO_TEXT)
|
||||
assert result is None
|
||||
|
||||
def test_get_using_provider_tts_disabled_returns_none(self, manager, mock_acm):
|
||||
mock_acm.get_conf.return_value = {
|
||||
"provider_tts_settings": {"enable": False}
|
||||
}
|
||||
result = manager.get_using_provider(ProviderType.TEXT_TO_SPEECH)
|
||||
assert result is None
|
||||
|
||||
def test_get_using_provider_unknown_type(self, manager, mock_acm):
|
||||
with pytest.raises(ValueError, match="Unknown provider type"):
|
||||
manager.get_using_provider(ProviderType.EMBEDDING)
|
||||
|
||||
def test_get_using_provider_with_umo(self, manager, mock_acm):
|
||||
fake_provider = MagicMock(spec=Provider)
|
||||
manager.inst_map["umo_prov"] = fake_provider
|
||||
mock_acm.get_conf.return_value = {
|
||||
"provider_settings": {"default_provider_id": None}
|
||||
}
|
||||
with patch("astrbot.core.provider.manager.sp") as mock_sp:
|
||||
mock_sp.get.return_value = "umo_prov"
|
||||
result = manager.get_using_provider(
|
||||
ProviderType.CHAT_COMPLETION, umo="session_1"
|
||||
)
|
||||
assert result is fake_provider
|
||||
mock_sp.get.assert_called_once()
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# _resolve_env_key_list
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestResolveEnvKeyList:
|
||||
def test_no_env_vars(self, manager):
|
||||
config = {"key": ["sk-abc", "sk-def"]}
|
||||
result = manager._resolve_env_key_list(config)
|
||||
assert result["key"] == ["sk-abc", "sk-def"]
|
||||
|
||||
def test_env_var_resolved(self, manager):
|
||||
os.environ["MY_API_KEY"] = "sk-from-env"
|
||||
config = {"key": ["$MY_API_KEY"], "id": "prov1"}
|
||||
result = manager._resolve_env_key_list(config)
|
||||
assert result["key"] == ["sk-from-env"]
|
||||
os.environ.pop("MY_API_KEY", None)
|
||||
|
||||
def test_env_var_braces(self, manager):
|
||||
os.environ["SECRET"] = "very_secret"
|
||||
config = {"key": ["${SECRET}"], "id": "prov1"}
|
||||
result = manager._resolve_env_key_list(config)
|
||||
assert result["key"] == ["very_secret"]
|
||||
os.environ.pop("SECRET", None)
|
||||
|
||||
def test_env_var_not_set_logs_warning(self, manager):
|
||||
# Ensure env var does not exist
|
||||
os.environ.pop("UNSET_VAR", None)
|
||||
config = {"key": ["$UNSET_VAR"], "id": "prov1"}
|
||||
with patch("astrbot.core.provider.manager.logger") as mock_logger:
|
||||
result = manager._resolve_env_key_list(config)
|
||||
assert result["key"] == [""]
|
||||
mock_logger.warning.assert_called_once()
|
||||
|
||||
def test_env_var_mixed_list(self, manager):
|
||||
os.environ["K1"] = "val1"
|
||||
config = {"key": ["$K1", "static_key"], "id": "prov1"}
|
||||
result = manager._resolve_env_key_list(config)
|
||||
assert result["key"] == ["val1", "static_key"]
|
||||
os.environ.pop("K1", None)
|
||||
|
||||
def test_non_list_key_passthrough(self, manager):
|
||||
config = {"key": "not_a_list"}
|
||||
result = manager._resolve_env_key_list(config)
|
||||
assert result["key"] == "not_a_list"
|
||||
|
||||
def test_empty_env_var_name(self, manager):
|
||||
config = {"key": ["$"], "id": "prov1"}
|
||||
result = manager._resolve_env_key_list(config)
|
||||
assert result["key"] == [""]
|
||||
|
||||
def test_missing_key_field(self, manager):
|
||||
config = {"id": "prov1"}
|
||||
result = manager._resolve_env_key_list(config)
|
||||
assert result == config
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# get_provider_config_by_id
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestGetProviderConfigById:
|
||||
def test_found(self, manager):
|
||||
manager.providers_config = [
|
||||
{"id": "p1", "type": "openai"},
|
||||
{"id": "p2", "type": "anthropic"},
|
||||
]
|
||||
result = manager.get_provider_config_by_id("p1")
|
||||
assert result == {"id": "p1", "type": "openai"}
|
||||
|
||||
def test_not_found(self, manager):
|
||||
manager.providers_config = [{"id": "p1", "type": "openai"}]
|
||||
result = manager.get_provider_config_by_id("nonexistent")
|
||||
assert result is None
|
||||
|
||||
def test_deep_copy_returned(self, manager):
|
||||
manager.providers_config = [{"id": "p1", "type": "openai", "key": ["secret"]}]
|
||||
result = manager.get_provider_config_by_id("p1")
|
||||
# Mutating the result should not affect the source
|
||||
result["key"].append("new_key")
|
||||
assert manager.providers_config[0]["key"] == ["secret"]
|
||||
|
||||
def test_merged_flag(self, manager):
|
||||
manager.providers_config = [
|
||||
{"id": "p1", "type": "openai", "provider_source_id": "src1"}
|
||||
]
|
||||
manager.provider_sources_config = [
|
||||
{"id": "src1", "base_url": "https://api.openai.com"}
|
||||
]
|
||||
with patch.object(
|
||||
manager,
|
||||
"get_merged_provider_config",
|
||||
return_value={"id": "p1", "type": "openai", "base_url": "https://api.openai.com"},
|
||||
):
|
||||
result = manager.get_provider_config_by_id("p1", merged=True)
|
||||
assert result["base_url"] == "https://api.openai.com"
|
||||
|
||||
def test_empty_configs(self, manager):
|
||||
result = manager.get_provider_config_by_id("p1")
|
||||
assert result is None
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# _get_all_provider_instances / _clear_loaded_instances / get_insts
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestProviderManagerInstances:
|
||||
def test_get_insts(self, manager):
|
||||
fp = MagicMock(spec=Provider)
|
||||
manager.provider_insts = [fp]
|
||||
assert manager.get_insts() == [fp]
|
||||
|
||||
def test_get_all_provider_instances_deduplicates(self, manager):
|
||||
fp = MagicMock(spec=Provider)
|
||||
manager.provider_insts = [fp]
|
||||
manager.inst_map = {"p1": fp}
|
||||
all_insts = manager._get_all_provider_instances()
|
||||
# fp appears in both lists but should only be returned once
|
||||
assert len(all_insts) == 1
|
||||
assert all_insts[0] is fp
|
||||
|
||||
def test_get_all_provider_instances_returns_all_types(self, manager):
|
||||
fp = MagicMock(spec=Provider)
|
||||
stt = MagicMock(spec=STTProvider)
|
||||
tts = MagicMock(spec=TTSProvider)
|
||||
emb = MagicMock(spec=EmbeddingProvider)
|
||||
rerank = MagicMock(spec=RerankProvider)
|
||||
manager.provider_insts = [fp]
|
||||
manager.stt_provider_insts = [stt]
|
||||
manager.tts_provider_insts = [tts]
|
||||
manager.embedding_provider_insts = [emb]
|
||||
manager.rerank_provider_insts = [rerank]
|
||||
all_insts = manager._get_all_provider_instances()
|
||||
assert len(all_insts) == 5
|
||||
|
||||
def test_clear_loaded_instances(self, manager):
|
||||
manager.provider_insts = [MagicMock(spec=Provider)]
|
||||
manager.stt_provider_insts = [MagicMock(spec=STTProvider)]
|
||||
manager.inst_map = {"p1": MagicMock()}
|
||||
manager.curr_provider_inst = MagicMock()
|
||||
|
||||
manager._clear_loaded_instances()
|
||||
assert manager.provider_insts == []
|
||||
assert manager.stt_provider_insts == []
|
||||
assert manager.tts_provider_insts == []
|
||||
assert manager.embedding_provider_insts == []
|
||||
assert manager.rerank_provider_insts == []
|
||||
assert manager.inst_map == {}
|
||||
assert manager.curr_provider_inst is None
|
||||
assert manager.curr_stt_provider_inst is None
|
||||
assert manager.curr_tts_provider_inst is None
|
||||
@@ -0,0 +1,473 @@
|
||||
"""Tests for astrbot.core.provider.provider.
|
||||
|
||||
Covers AbstractProvider, Provider, STTProvider, TTSProvider,
|
||||
EmbeddingProvider, and RerankProvider abstract/concrete methods.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.agent.message import Message
|
||||
from astrbot.core.provider.entities import (
|
||||
LLMResponse,
|
||||
ProviderMeta,
|
||||
ProviderMetaData,
|
||||
ProviderType,
|
||||
)
|
||||
from astrbot.core.provider.provider import (
|
||||
AbstractProvider,
|
||||
EmbeddingProvider,
|
||||
Provider,
|
||||
RerankProvider,
|
||||
STTProvider,
|
||||
TTSProvider,
|
||||
)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# AbstractProvider
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestAbstractProvider:
|
||||
def test_construction(self):
|
||||
ap = AbstractProvider(provider_config={"type": "test"})
|
||||
assert ap.model_name == ""
|
||||
assert ap.provider_config == {"type": "test"}
|
||||
|
||||
def test_set_and_get_model(self):
|
||||
ap = AbstractProvider(provider_config={"type": "test"})
|
||||
assert ap.get_model() == ""
|
||||
ap.set_model("gpt-4")
|
||||
assert ap.get_model() == "gpt-4"
|
||||
|
||||
def test_set_model_empty_string(self):
|
||||
ap = AbstractProvider(provider_config={"type": "test"})
|
||||
ap.set_model("gpt-4")
|
||||
ap.set_model("")
|
||||
assert ap.get_model() == ""
|
||||
|
||||
def test_meta_returns_provider_meta(self):
|
||||
pmd = ProviderMetaData(
|
||||
id="default",
|
||||
model=None,
|
||||
type="test_type",
|
||||
provider_type=ProviderType.CHAT_COMPLETION,
|
||||
)
|
||||
with patch(
|
||||
"astrbot.core.provider.provider.provider_cls_map",
|
||||
{"test_type": pmd},
|
||||
):
|
||||
ap = AbstractProvider(provider_config={"type": "test_type", "id": "myid"})
|
||||
meta = ap.meta()
|
||||
assert isinstance(meta, ProviderMeta)
|
||||
assert meta.id == "myid"
|
||||
assert meta.type == "test_type"
|
||||
assert meta.provider_type == ProviderType.CHAT_COMPLETION
|
||||
|
||||
def test_meta_raises_on_unregistered_type(self):
|
||||
ap = AbstractProvider(provider_config={"type": "nonexistent"})
|
||||
with pytest.raises(ValueError, match="not registered"):
|
||||
ap.meta()
|
||||
|
||||
def test_meta_no_provider_config_id_falls_back_to_default(self):
|
||||
pmd = ProviderMetaData(
|
||||
id="default", model=None, type="test_type", provider_type=ProviderType.EMBEDDING
|
||||
)
|
||||
with patch(
|
||||
"astrbot.core.provider.provider.provider_cls_map",
|
||||
{"test_type": pmd},
|
||||
):
|
||||
ap = AbstractProvider(provider_config={"type": "test_type"})
|
||||
meta = ap.meta()
|
||||
assert meta.id == "default"
|
||||
|
||||
def test_test_does_not_raise(self):
|
||||
ap = AbstractProvider(provider_config={"type": "test"})
|
||||
# test() is a no-op on AbstractProvider
|
||||
ap.test() # should not raise
|
||||
|
||||
def test_constructor_sets_provider_config(self):
|
||||
config = {"type": "myprovider", "key": ["sk-abc"]}
|
||||
ap = AbstractProvider(provider_config=config)
|
||||
assert ap.provider_config is config
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Provider (Chat)
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class _ConcreteProvider(Provider):
|
||||
"""Minimal concrete subclass for testing Provider abstract methods."""
|
||||
|
||||
def get_current_key(self) -> str:
|
||||
return "key_override"
|
||||
|
||||
def set_key(self, key: str) -> None:
|
||||
self._key = key
|
||||
|
||||
async def get_models(self) -> list[str]:
|
||||
return ["model-a", "model-b"]
|
||||
|
||||
async def text_chat(self, **kwargs) -> LLMResponse:
|
||||
return LLMResponse(role="assistant", completion_text="mock reply")
|
||||
|
||||
|
||||
class TestProvider:
|
||||
def test_construction(self):
|
||||
p = _ConcreteProvider(
|
||||
provider_config={"type": "test", "key": ["sk-abc"]},
|
||||
provider_settings={},
|
||||
)
|
||||
assert p.provider_config["type"] == "test"
|
||||
assert p.provider_settings == {}
|
||||
|
||||
def test_get_current_key_abstract(self):
|
||||
p = _ConcreteProvider(provider_config={"type": "test"}, provider_settings={})
|
||||
assert p.get_current_key() == "key_override"
|
||||
|
||||
def test_get_keys_default(self):
|
||||
p = _ConcreteProvider(provider_config={"type": "test"}, provider_settings={})
|
||||
assert p.get_keys() == [""]
|
||||
|
||||
def test_get_keys_from_config(self):
|
||||
p = _ConcreteProvider(
|
||||
provider_config={"type": "test", "key": ["sk-1", "sk-2"]},
|
||||
provider_settings={},
|
||||
)
|
||||
assert p.get_keys() == ["sk-1", "sk-2"]
|
||||
|
||||
def test_get_keys_none(self):
|
||||
p = _ConcreteProvider(
|
||||
provider_config={"type": "test", "key": None},
|
||||
provider_settings={},
|
||||
)
|
||||
assert p.get_keys() == [""]
|
||||
|
||||
def test_set_key(self):
|
||||
p = _ConcreteProvider(provider_config={"type": "test"}, provider_settings={})
|
||||
p.set_key("new-key")
|
||||
assert p._key == "new-key"
|
||||
|
||||
def test_get_models(self):
|
||||
p = _ConcreteProvider(provider_config={"type": "test"}, provider_settings={})
|
||||
models = p.get_models()
|
||||
assert models == ["model-a", "model-b"]
|
||||
|
||||
def test_text_chat(self):
|
||||
p = _ConcreteProvider(provider_config={"type": "test"}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
resp = asyncio.run(p.text_chat(prompt="hi"))
|
||||
assert isinstance(resp, LLMResponse)
|
||||
assert resp.role == "assistant"
|
||||
assert resp.completion_text == "mock reply"
|
||||
|
||||
def test_text_chat_abstract_prevents_instantiation(self):
|
||||
"""Provider subclasses must override text_chat; can't instantiate without it."""
|
||||
class IncompleteProvider(Provider):
|
||||
def get_current_key(self) -> str:
|
||||
return ""
|
||||
|
||||
def set_key(self, key: str) -> None:
|
||||
pass
|
||||
|
||||
async def get_models(self) -> list[str]:
|
||||
return []
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
IncompleteProvider(provider_config={}, provider_settings={})
|
||||
|
||||
def test_text_chat_stream_raises_not_implemented(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
with pytest.raises(NotImplementedError):
|
||||
import asyncio
|
||||
|
||||
async def consume():
|
||||
gen = p.text_chat_stream(prompt="hi")
|
||||
async for _ in gen:
|
||||
pass
|
||||
|
||||
asyncio.run(consume())
|
||||
|
||||
def test_pop_record_removes_non_system(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
ctx = [
|
||||
{"role": "system", "content": "You are a bot"},
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "hi"},
|
||||
{"role": "user", "content": "how are you"},
|
||||
]
|
||||
p.pop_record(ctx)
|
||||
assert len(ctx) == 2
|
||||
assert ctx[0]["role"] == "system"
|
||||
assert ctx[1]["role"] == "user"
|
||||
assert ctx[1]["content"] == "how are you"
|
||||
|
||||
def test_pop_record_with_no_system(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
ctx = [
|
||||
{"role": "user", "content": "first"},
|
||||
{"role": "assistant", "content": "reply"},
|
||||
]
|
||||
p.pop_record(ctx)
|
||||
# Both should be popped
|
||||
assert len(ctx) == 0
|
||||
|
||||
def test_pop_record_capped_at_two(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
ctx = [
|
||||
{"role": "system", "content": "sys"},
|
||||
{"role": "user", "content": "u1"},
|
||||
{"role": "assistant", "content": "a1"},
|
||||
{"role": "user", "content": "u2"},
|
||||
{"role": "assistant", "content": "a2"},
|
||||
{"role": "user", "content": "u3"},
|
||||
]
|
||||
p.pop_record(ctx)
|
||||
assert len(ctx) == 4
|
||||
# system kept + last 3 non-system (u2, a2, u3) → but pop_record removes first 2 non-system
|
||||
assert ctx[0] == {"role": "system", "content": "sys"}
|
||||
assert ctx[1] == {"role": "user", "content": "u2"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_test(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
with patch.object(p, "text_chat", AsyncMock(return_value=LLMResponse(role="assistant"))):
|
||||
await p.test(test_timeout=5.0)
|
||||
|
||||
def test_ensure_message_to_dicts_none(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
result = p._ensure_message_to_dicts(None)
|
||||
assert result == []
|
||||
|
||||
def test_ensure_message_to_dicts_empty(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
result = p._ensure_message_to_dicts([])
|
||||
assert result == []
|
||||
|
||||
def test_ensure_message_to_dicts_skips_checkpoint(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
checkpoint = {"role": "user", "content": "check", "checkpoint": True}
|
||||
normal = {"role": "user", "content": "normal"}
|
||||
with patch(
|
||||
"astrbot.core.provider.provider.is_checkpoint_message",
|
||||
side_effect=lambda m: m.get("checkpoint", False),
|
||||
):
|
||||
result = p._ensure_message_to_dicts([checkpoint, normal])
|
||||
assert len(result) == 1
|
||||
assert result[0] == {"role": "user", "content": "normal"}
|
||||
|
||||
def test_ensure_message_to_dicts_converts_message_objects(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
msg = MagicMock(spec=Message)
|
||||
msg.model_dump.return_value = {"role": "user", "content": "from pydantic"}
|
||||
result = p._ensure_message_to_dicts([msg])
|
||||
assert result == [{"role": "user", "content": "from pydantic"}]
|
||||
|
||||
def test_ensure_message_to_dicts_passes_through_dicts(self):
|
||||
p = _ConcreteProvider(provider_config={}, provider_settings={})
|
||||
d = {"role": "user", "content": "plain dict"}
|
||||
result = p._ensure_message_to_dicts([d])
|
||||
assert result == [d]
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# STTProvider
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class _ConcreteSTTProvider(STTProvider):
|
||||
async def get_text(self, audio_url: str) -> str:
|
||||
return "transcribed text"
|
||||
|
||||
|
||||
class TestSTTProvider:
|
||||
def test_construction(self):
|
||||
p = _ConcreteSTTProvider(
|
||||
provider_config={"type": "stt_test"},
|
||||
provider_settings={},
|
||||
)
|
||||
assert p.provider_config["type"] == "stt_test"
|
||||
|
||||
def test_get_text(self):
|
||||
p = _ConcreteSTTProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
text = asyncio.run(p.get_text("/path/to/audio.wav"))
|
||||
assert text == "transcribed text"
|
||||
|
||||
def test_get_text_abstract(self):
|
||||
class IncompleteSTT(STTProvider):
|
||||
pass
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
IncompleteSTT(provider_config={}, provider_settings={})
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# TTSProvider
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class _ConcreteTTSProvider(TTSProvider):
|
||||
async def get_audio(self, text: str) -> str:
|
||||
return "/tmp/test_output.wav"
|
||||
|
||||
|
||||
class TestTTSProvider:
|
||||
def test_construction(self):
|
||||
p = _ConcreteTTSProvider(provider_config={"type": "tts_test"}, provider_settings={})
|
||||
assert p.provider_config["type"] == "tts_test"
|
||||
|
||||
def test_get_audio(self):
|
||||
p = _ConcreteTTSProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
path = asyncio.run(p.get_audio("hello"))
|
||||
assert path == "/tmp/test_output.wav"
|
||||
|
||||
def test_support_stream_default(self):
|
||||
p = _ConcreteTTSProvider(provider_config={}, provider_settings={})
|
||||
assert p.support_stream() is False
|
||||
|
||||
def test_get_audio_stream_default_implementation(self):
|
||||
p = _ConcreteTTSProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
text_queue: asyncio.Queue[str | None] = asyncio.Queue()
|
||||
audio_queue: asyncio.Queue = asyncio.Queue()
|
||||
|
||||
async def run_stream():
|
||||
# Send some text, then None to signal end
|
||||
await text_queue.put("hello ")
|
||||
await text_queue.put("world")
|
||||
await text_queue.put(None)
|
||||
await p.get_audio_stream(text_queue, audio_queue)
|
||||
|
||||
asyncio.run(run_stream())
|
||||
result = asyncio.run(audio_queue.get())
|
||||
assert result is not None
|
||||
text_part, audio_data = result
|
||||
assert text_part == "hello world"
|
||||
assert isinstance(audio_data, bytes)
|
||||
# The None sentinel should follow
|
||||
end = asyncio.run(audio_queue.get())
|
||||
assert end is None
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# EmbeddingProvider
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class _ConcreteEmbeddingProvider(EmbeddingProvider):
|
||||
async def get_embedding(self, text: str) -> list[float]:
|
||||
return [0.1, 0.2, 0.3]
|
||||
|
||||
async def get_embeddings(self, text: list[str]) -> list[list[float]]:
|
||||
return [[0.1, 0.2, 0.3] for _ in text]
|
||||
|
||||
def get_dim(self) -> int:
|
||||
return 3
|
||||
|
||||
|
||||
class TestEmbeddingProvider:
|
||||
def test_construction(self):
|
||||
p = _ConcreteEmbeddingProvider(
|
||||
provider_config={"type": "emb_test"},
|
||||
provider_settings={},
|
||||
)
|
||||
assert p.provider_config["type"] == "emb_test"
|
||||
|
||||
def test_get_embedding(self):
|
||||
p = _ConcreteEmbeddingProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
emb = asyncio.run(p.get_embedding("hello"))
|
||||
assert emb == [0.1, 0.2, 0.3]
|
||||
|
||||
def test_get_embeddings(self):
|
||||
p = _ConcreteEmbeddingProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
embs = asyncio.run(p.get_embeddings(["a", "b"]))
|
||||
assert len(embs) == 2
|
||||
|
||||
def test_get_dim(self):
|
||||
p = _ConcreteEmbeddingProvider(provider_config={}, provider_settings={})
|
||||
assert p.get_dim() == 3
|
||||
|
||||
def test_get_embeddings_batch_single_batch(self):
|
||||
p = _ConcreteEmbeddingProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
embs = asyncio.run(p.get_embeddings_batch(["hello", "world"], batch_size=10))
|
||||
assert len(embs) == 2
|
||||
|
||||
def test_get_embeddings_batch_multiple_batches(self):
|
||||
p = _ConcreteEmbeddingProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
texts = [f"text_{i}" for i in range(5)]
|
||||
embs = asyncio.run(p.get_embeddings_batch(texts, batch_size=2, tasks_limit=5))
|
||||
assert len(embs) == 5
|
||||
|
||||
def test_get_embeddings_batch_with_progress_callback(self):
|
||||
p = _ConcreteEmbeddingProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
progress = AsyncMock()
|
||||
texts = [f"t{i}" for i in range(3)]
|
||||
embs = asyncio.run(
|
||||
p.get_embeddings_batch(texts, batch_size=2, tasks_limit=5, progress_callback=progress)
|
||||
)
|
||||
assert len(embs) == 3
|
||||
assert progress.await_count >= 1
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# RerankProvider
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class _ConcreteRerankProvider(RerankProvider):
|
||||
async def rerank(
|
||||
self,
|
||||
query: str,
|
||||
documents: list[str],
|
||||
top_n: int | None = None,
|
||||
):
|
||||
from astrbot.core.provider.entities import RerankResult
|
||||
|
||||
return [
|
||||
RerankResult(index=0, relevance_score=0.95),
|
||||
RerankResult(index=1, relevance_score=0.80),
|
||||
][:top_n]
|
||||
|
||||
|
||||
class TestRerankProvider:
|
||||
def test_construction(self):
|
||||
p = _ConcreteRerankProvider(
|
||||
provider_config={"type": "rerank_test"},
|
||||
provider_settings={},
|
||||
)
|
||||
assert p.provider_config["type"] == "rerank_test"
|
||||
|
||||
def test_rerank(self):
|
||||
p = _ConcreteRerankProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
results = asyncio.run(p.rerank("test query", ["doc1", "doc2"]))
|
||||
assert len(results) == 2
|
||||
assert results[0].index == 0
|
||||
assert results[0].relevance_score == 0.95
|
||||
|
||||
def test_rerank_with_top_n(self):
|
||||
p = _ConcreteRerankProvider(provider_config={}, provider_settings={})
|
||||
import asyncio
|
||||
|
||||
results = asyncio.run(p.rerank("test query", ["doc1", "doc2", "doc3"], top_n=1))
|
||||
assert len(results) == 1
|
||||
@@ -0,0 +1,244 @@
|
||||
"""Tests for astrbot.core.provider.register.
|
||||
|
||||
Covers register_provider_adapter decorator, provider_registry list,
|
||||
provider_cls_map dict, llm_tools import, duplicate registration,
|
||||
and default config template processing.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.provider.entities import ProviderMetaData, ProviderType
|
||||
from astrbot.core.provider.register import (
|
||||
llm_tools,
|
||||
provider_cls_map,
|
||||
provider_registry,
|
||||
register_provider_adapter,
|
||||
)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Fixtures — fresh state per test
|
||||
# =========================================================================
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_registries():
|
||||
"""Clear global registries before and after each test to avoid cross-test pollution."""
|
||||
before_keys = set(provider_cls_map.keys())
|
||||
before_len = len(provider_registry)
|
||||
|
||||
yield
|
||||
|
||||
# Restore: remove any keys/items added during the test
|
||||
added_keys = set(provider_cls_map.keys()) - before_keys
|
||||
for k in added_keys:
|
||||
provider_cls_map.pop(k, None)
|
||||
|
||||
added_items = len(provider_registry) - before_len
|
||||
for _ in range(added_items):
|
||||
if provider_registry:
|
||||
provider_registry.pop()
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# register_provider_adapter
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestRegisterProviderAdapter:
|
||||
def test_basic_registration(self):
|
||||
@register_provider_adapter("test_provider", "A test provider")
|
||||
class FakeProvider:
|
||||
pass
|
||||
|
||||
assert "test_provider" in provider_cls_map
|
||||
pmd = provider_cls_map["test_provider"]
|
||||
assert isinstance(pmd, ProviderMetaData)
|
||||
assert pmd.type == "test_provider"
|
||||
assert pmd.desc == "A test provider"
|
||||
assert pmd.cls_type is FakeProvider
|
||||
assert pmd.provider_type == ProviderType.CHAT_COMPLETION # default
|
||||
# Should also be in the registry list
|
||||
assert pmd in provider_registry
|
||||
|
||||
def test_registration_with_custom_provider_type(self):
|
||||
@register_provider_adapter(
|
||||
"stt_provider",
|
||||
"STT provider",
|
||||
provider_type=ProviderType.SPEECH_TO_TEXT,
|
||||
)
|
||||
class FakeSTT:
|
||||
pass
|
||||
|
||||
pmd = provider_cls_map["stt_provider"]
|
||||
assert pmd.provider_type == ProviderType.SPEECH_TO_TEXT
|
||||
|
||||
def test_registration_with_display_name(self):
|
||||
@register_provider_adapter(
|
||||
"disp_provider",
|
||||
"Display test",
|
||||
provider_display_name="My Display Name",
|
||||
)
|
||||
class DispProvider:
|
||||
pass
|
||||
|
||||
pmd = provider_cls_map["disp_provider"]
|
||||
assert pmd.provider_display_name == "My Display Name"
|
||||
|
||||
def test_registration_with_default_config_tmpl(self):
|
||||
@register_provider_adapter(
|
||||
"tmpl_provider",
|
||||
"Template test",
|
||||
default_config_tmpl={"key1": "val1"},
|
||||
)
|
||||
class TmplProvider:
|
||||
pass
|
||||
|
||||
pmd = provider_cls_map["tmpl_provider"]
|
||||
assert pmd.default_config_tmpl is not None
|
||||
# The decorator adds mandatory fields
|
||||
assert pmd.default_config_tmpl["type"] == "tmpl_provider"
|
||||
assert pmd.default_config_tmpl["enable"] is False
|
||||
assert pmd.default_config_tmpl["id"] == "tmpl_provider"
|
||||
assert pmd.default_config_tmpl["key1"] == "val1"
|
||||
|
||||
def test_default_config_tmpl_preserves_existing_type(self):
|
||||
@register_provider_adapter(
|
||||
"tmpl2",
|
||||
"desc",
|
||||
default_config_tmpl={"type": "custom_type", "custom": True},
|
||||
)
|
||||
class T2:
|
||||
pass
|
||||
|
||||
pmd = provider_cls_map["tmpl2"]
|
||||
# should NOT override existing type
|
||||
assert pmd.default_config_tmpl["type"] == "custom_type"
|
||||
assert pmd.default_config_tmpl["enable"] is False
|
||||
assert pmd.default_config_tmpl["id"] == "tmpl2"
|
||||
assert pmd.default_config_tmpl["custom"] is True
|
||||
|
||||
def test_default_config_tmpl_preserves_existing_enable(self):
|
||||
@register_provider_adapter(
|
||||
"tmpl3", "desc", default_config_tmpl={"enable": True, "extra": "x"}
|
||||
)
|
||||
class T3:
|
||||
pass
|
||||
|
||||
pmd = provider_cls_map["tmpl3"]
|
||||
assert pmd.default_config_tmpl["enable"] is True
|
||||
assert pmd.default_config_tmpl["type"] == "tmpl3"
|
||||
assert pmd.default_config_tmpl["extra"] == "x"
|
||||
|
||||
def test_default_config_tmpl_preserves_existing_id(self):
|
||||
@register_provider_adapter(
|
||||
"tmpl4",
|
||||
"desc",
|
||||
default_config_tmpl={"id": "custom_id"},
|
||||
)
|
||||
class T4:
|
||||
pass
|
||||
|
||||
pmd = provider_cls_map["tmpl4"]
|
||||
assert pmd.default_config_tmpl["id"] == "custom_id"
|
||||
|
||||
def test_default_config_tmpl_none(self):
|
||||
@register_provider_adapter("no_tmpl", "No template")
|
||||
class NoTmpl:
|
||||
pass
|
||||
|
||||
pmd = provider_cls_map["no_tmpl"]
|
||||
assert pmd.default_config_tmpl is None
|
||||
|
||||
def test_duplicate_registration_raises(self):
|
||||
@register_provider_adapter("dup_provider", "First")
|
||||
class FirstProvider:
|
||||
pass
|
||||
|
||||
with pytest.raises(ValueError, match="已经注册"):
|
||||
@register_provider_adapter("dup_provider", "Second")
|
||||
class SecondProvider:
|
||||
pass
|
||||
|
||||
# The original registration should remain intact
|
||||
assert provider_cls_map["dup_provider"].desc == "First"
|
||||
assert provider_cls_map["dup_provider"].cls_type is FirstProvider
|
||||
|
||||
def test_multiple_registrations(self):
|
||||
@register_provider_adapter("p1", "First provider")
|
||||
class P1:
|
||||
pass
|
||||
|
||||
@register_provider_adapter("p2", "Second provider")
|
||||
class P2:
|
||||
pass
|
||||
|
||||
assert len(provider_cls_map) == 2
|
||||
assert provider_cls_map["p1"].cls_type is P1
|
||||
assert provider_cls_map["p2"].cls_type is P2
|
||||
assert len(provider_registry) >= 2
|
||||
|
||||
def test_decorator_returns_the_class(self):
|
||||
@register_provider_adapter("return_test", "Check return")
|
||||
class ReturnClass:
|
||||
pass
|
||||
|
||||
# The decorator should return the class unchanged so it can be used normally
|
||||
instance = ReturnClass()
|
||||
assert isinstance(instance, ReturnClass)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# llm_tools / FuncCall
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestLlmTools:
|
||||
def test_llm_tools_is_importable(self):
|
||||
from astrbot.core.provider.register import llm_tools as lt
|
||||
|
||||
assert lt is not None
|
||||
|
||||
def test_llm_tools_function_tool_manager_type(self):
|
||||
from astrbot.core.provider.func_tool_manager import FunctionToolManager
|
||||
|
||||
assert isinstance(llm_tools, FunctionToolManager)
|
||||
|
||||
def test_llm_tools_starts_empty(self):
|
||||
assert llm_tools.empty() is True
|
||||
|
||||
def test_llm_tools_add_and_get(self):
|
||||
async def fake_handler(**kwargs):
|
||||
return "done"
|
||||
|
||||
llm_tools.add_func(
|
||||
name="test_func",
|
||||
func_args=[{"name": "arg1", "type": "string", "description": "An arg"}],
|
||||
desc="A test function",
|
||||
handler=fake_handler,
|
||||
)
|
||||
tool = llm_tools.get_func("test_func")
|
||||
assert tool is not None
|
||||
assert tool.name == "test_func"
|
||||
llm_tools.remove_func("test_func")
|
||||
assert llm_tools.get_func("test_func") is None
|
||||
|
||||
def test_llm_tools_remove_nonexistent(self):
|
||||
"""Removing a function that does not exist should not raise."""
|
||||
llm_tools.remove_func("nonexistent_tool") # should not raise
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Registry list / map invariants
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestRegistryInvariants:
|
||||
def test_meta_id_is_default_in_registry(self):
|
||||
@register_provider_adapter("invariant_test", "check id")
|
||||
class InvProvider:
|
||||
pass
|
||||
|
||||
pmd = provider_cls_map["invariant_test"]
|
||||
assert pmd.id == "default"
|
||||
assert pmd.model is None
|
||||
@@ -0,0 +1,375 @@
|
||||
"""
|
||||
Unit tests for RecursiveCharacterChunker.
|
||||
|
||||
Covers construction, chunk method with various inputs (empty text, short text,
|
||||
newline separators, character-level fallback, recursive splitting, custom
|
||||
separators, and edge cases with chunk_size/overlap validation).
|
||||
All tests isolate the chunker from any external dependencies.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.knowledge_base.chunking.recursive import (
|
||||
RecursiveCharacterChunker,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Construction
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRecursiveCharacterChunkerConstruction:
|
||||
"""Test construction of RecursiveCharacterChunker."""
|
||||
|
||||
def test_default_construction(self):
|
||||
"""Test default parameters are set correctly."""
|
||||
chunker = RecursiveCharacterChunker()
|
||||
assert chunker.chunk_size == 500
|
||||
assert chunker.chunk_overlap == 100
|
||||
assert chunker.length_function is len
|
||||
assert chunker.is_separator_regex is False
|
||||
assert "\n\n" in chunker.separators
|
||||
assert "" in chunker.separators # character fallback
|
||||
|
||||
def test_custom_construction(self):
|
||||
"""Test custom parameters are applied."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=1000,
|
||||
chunk_overlap=200,
|
||||
length_function=lambda x: len(x.split()),
|
||||
is_separator_regex=True,
|
||||
separators=["\n", " "],
|
||||
)
|
||||
assert chunker.chunk_size == 1000
|
||||
assert chunker.chunk_overlap == 200
|
||||
assert chunker.length_function("hello world") == 2
|
||||
assert chunker.is_separator_regex is True
|
||||
assert chunker.separators == ["\n", " "]
|
||||
|
||||
def test_injectable_length_function(self):
|
||||
"""Test injecting a custom length function."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=3,
|
||||
length_function=lambda x: len(x.split()),
|
||||
)
|
||||
assert chunker.chunk_size == 3
|
||||
assert chunker.length_function("a b c d") == 4
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - basic cases
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRecursiveCharacterChunkerChunkBasic:
|
||||
"""Test basic chunk() behavior."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_text_returns_empty_list(self):
|
||||
"""Test that empty text returns an empty list."""
|
||||
chunker = RecursiveCharacterChunker()
|
||||
result = await chunker.chunk("")
|
||||
assert result == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_whitespace_only_text(self):
|
||||
"""Test that whitespace-only text is handled."""
|
||||
chunker = RecursiveCharacterChunker()
|
||||
result = await chunker.chunk(" \n\n ")
|
||||
# Depends on separator logic; should not crash
|
||||
assert isinstance(result, list)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_text_returns_single_chunk(self):
|
||||
"""Test that text shorter than chunk_size returns single chunk."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=1000)
|
||||
text = "Short text."
|
||||
result = await chunker.chunk(text)
|
||||
assert result == [text]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_equal_to_chunk_size(self):
|
||||
"""Test that text exactly chunk_size returns single chunk."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=10)
|
||||
text = "0123456789"
|
||||
result = await chunker.chunk(text)
|
||||
assert result == [text]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - separator-driven splitting
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRecursiveCharacterChunkerSeparator:
|
||||
"""Test chunk() behavior with separator-based splitting."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_splits_by_double_newline(self):
|
||||
"""Test that double newline is the preferred separator."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=10, chunk_overlap=0)
|
||||
text = "para1\n\npara2\n\npara3"
|
||||
result = await chunker.chunk(text)
|
||||
# Each paragraph is <= chunk_size, so each should be a separate chunk
|
||||
assert len(result) >= 1
|
||||
# All chunks should be non-empty strings
|
||||
assert all(isinstance(c, str) and c for c in result)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_splits_by_newline_when_double_newline_not_found(self):
|
||||
"""Test fallback to single newline when double newline is absent."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=50,
|
||||
chunk_overlap=0,
|
||||
separators=["\n\n", "\n"],
|
||||
)
|
||||
text = "line1\nline2\nline3"
|
||||
result = await chunker.chunk(text)
|
||||
# Without overlap and small lines, each line should be separate if lines fit
|
||||
assert isinstance(result, list)
|
||||
assert all(isinstance(c, str) for c in result)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_falls_back_to_character_splitting(self):
|
||||
"""Test that chunker falls back to character splitting when no separator matches."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=5,
|
||||
chunk_overlap=0,
|
||||
separators=["\n\n", ""],
|
||||
)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
# Character splitting: step = chunk_size - overlap = 5
|
||||
assert result == ["abcde", "fghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_separator(self):
|
||||
"""Test that a custom separator is used for splitting."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=100,
|
||||
chunk_overlap=0,
|
||||
separators=["|"],
|
||||
)
|
||||
text = "part1|part2|part3"
|
||||
result = await chunker.chunk(text)
|
||||
# Each part is well under chunk_size, so each should be a separate chunk
|
||||
assert len(result) == 3
|
||||
assert result[0] == "part1|"
|
||||
assert result[1] == "part2|"
|
||||
assert result[2] == "part3"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - overlap behavior
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRecursiveCharacterChunkerOverlap:
|
||||
"""Test overlap behavior in chunk()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_character_split_with_overlap(self):
|
||||
"""Test that character-level split respects overlap."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=6,
|
||||
chunk_overlap=2,
|
||||
separators=[""],
|
||||
)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
# step = 6 - 2 = 4, so: "abcdef", "efghij"
|
||||
assert result == ["abcdef", "efghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_overlap_from_kwargs_overrides_instance_default(self):
|
||||
"""Test that kwargs chunk_overlap overrides the instance default."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=6,
|
||||
chunk_overlap=0,
|
||||
separators=[""],
|
||||
)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text, chunk_overlap=2)
|
||||
# step = 6 - 2 = 4
|
||||
assert result == ["abcdef", "efghij"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chunk_size_from_kwargs_overrides_instance_default(self):
|
||||
"""Test that kwargs chunk_size overrides the instance default."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=500,
|
||||
chunk_overlap=0,
|
||||
separators=[""],
|
||||
)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text, chunk_size=5)
|
||||
# step = 5 - 0 = 5
|
||||
assert result == ["abcde", "fghij"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - recursive splitting of oversized segments
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRecursiveCharacterChunkerRecursive:
|
||||
"""Test recursive splitting of segments that exceed chunk_size."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recursive_split_oversized_segment(self):
|
||||
"""Test that a single segment larger than chunk_size is recursively split."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=10,
|
||||
chunk_overlap=0,
|
||||
separators=["\n\n", ""],
|
||||
)
|
||||
# Single paragraph (no double newline) that exceeds chunk_size
|
||||
text = "a" * 25
|
||||
result = await chunker.chunk(text)
|
||||
# Should split into 3 character-level chunks: 10 + 10 + 5
|
||||
assert len(result) == 3
|
||||
assert result[0] == "a" * 10
|
||||
assert result[1] == "a" * 10
|
||||
assert result[2] == "a" * 5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recursive_and_normal_chunks_mixed(self):
|
||||
"""Test mixing normal chunks and recursively split oversized segments."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=10,
|
||||
chunk_overlap=0,
|
||||
separators=["\n\n", ""],
|
||||
)
|
||||
# First paragraph fits, second paragraph is oversized, third fits
|
||||
text = "short" + "\n\n" + ("a" * 25) + "\n\n" + "tiny"
|
||||
result = await chunker.chunk(text)
|
||||
# short\n\n + a*10 + a*10 + a*5 + \n\n + tiny -> should be 4 chunks
|
||||
# Note: the separator is included in each split
|
||||
assert len(result) >= 3
|
||||
assert all(isinstance(c, str) and c for c in result)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# _split_by_character
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRecursiveCharacterChunkerSplitByCharacter:
|
||||
"""Test the _split_by_character private method."""
|
||||
|
||||
def test_split_by_character_defaults(self):
|
||||
"""Test _split_by_character uses instance defaults."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=4, chunk_overlap=1)
|
||||
result = chunker._split_by_character("abcdefgh")
|
||||
assert result == ["abcd", "defg", "gh"]
|
||||
|
||||
def test_split_by_character_explicit_params(self):
|
||||
"""Test _split_by_character with explicit chunk_size and overlap."""
|
||||
chunker = RecursiveCharacterChunker()
|
||||
result = chunker._split_by_character("abcdefgh", chunk_size=4, overlap=1)
|
||||
assert result == ["abcd", "defg", "gh"]
|
||||
|
||||
def test_split_by_character_long_exact_fit(self):
|
||||
"""Test split when text length divides evenly into chunk_size."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=4, chunk_overlap=0)
|
||||
result = chunker._split_by_character("abcdefgh")
|
||||
assert result == ["abcd", "efgh"]
|
||||
|
||||
def test_split_by_character_short_text(self):
|
||||
"""Test split when text is shorter than chunk_size."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=100)
|
||||
result = chunker._split_by_character("hello")
|
||||
assert result == ["hello"]
|
||||
|
||||
def test_split_by_character_raises_on_zero_chunk_size(self):
|
||||
"""Test that chunk_size <= 0 raises ValueError."""
|
||||
chunker = RecursiveCharacterChunker()
|
||||
with pytest.raises(ValueError, match="chunk_size must be greater than 0"):
|
||||
chunker._split_by_character("test", chunk_size=0)
|
||||
|
||||
def test_split_by_character_raises_on_negative_overlap(self):
|
||||
"""Test that negative overlap raises ValueError."""
|
||||
chunker = RecursiveCharacterChunker()
|
||||
with pytest.raises(ValueError, match="chunk_overlap must be non-negative"):
|
||||
chunker._split_by_character("test", chunk_overlap=-1)
|
||||
|
||||
def test_split_by_character_raises_on_overlap_ge_chunk_size(self):
|
||||
"""Test that overlap >= chunk_size raises ValueError."""
|
||||
chunker = RecursiveCharacterChunker()
|
||||
with pytest.raises(ValueError, match="chunk_overlap must be less than chunk_size"):
|
||||
chunker._split_by_character("test", chunk_size=5, chunk_overlap=5)
|
||||
|
||||
def test_single_character_chunks(self):
|
||||
"""Test splitting with chunk_size=1."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=1, chunk_overlap=0)
|
||||
result = chunker._split_by_character("abc")
|
||||
assert result == ["a", "b", "c"]
|
||||
|
||||
def test_overlap_equal_chunk_size_minus_one(self):
|
||||
"""Test with maximum valid overlap."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=5, chunk_overlap=4)
|
||||
result = chunker._split_by_character("abcdefghij")
|
||||
# step = 1, each chunk slides by 1 char
|
||||
assert result == ["abcde", "bcdef", "cdefg", "defgh", "efghi", "fghij"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# chunk() - integration scenarios
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRecursiveCharacterChunkerIntegration:
|
||||
"""Integration-style tests combining multiple logic paths."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_paragraphs_with_oversized_lines(self):
|
||||
"""Test a realistic paragraph mix with some oversized lines."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=20, chunk_overlap=0)
|
||||
lines = ["A" * 30, "short", "B" * 30]
|
||||
text = "\n".join(lines)
|
||||
result = await chunker.chunk(text)
|
||||
# Should split the oversized lines but keep short lines as-is
|
||||
assert all(isinstance(c, str) and c for c in result)
|
||||
assert len(result) >= 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_separator_in_text_but_not_at_separator_boundary(self):
|
||||
"""Test that chunker finds the best separator when multiple are present."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=10,
|
||||
chunk_overlap=0,
|
||||
separators=["\n\n", "\n", ""],
|
||||
)
|
||||
# Has both double-newline and single-newline; prefer double-newline
|
||||
text = "aaaa\n\nbbbb\ncccc"
|
||||
result = await chunker.chunk(text)
|
||||
# Should split on \n\n first, then possibly \n for remaining
|
||||
assert isinstance(result, list)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_preserves_ordering(self):
|
||||
"""Test that chunks preserve the original text ordering."""
|
||||
chunker = RecursiveCharacterChunker(chunk_size=10, chunk_overlap=0)
|
||||
text = "first\n\nsecond\n\nthird\n\nfourth"
|
||||
result = await chunker.chunk(text)
|
||||
# Concatenation of all chunks should contain all words in order
|
||||
combined = "".join(result)
|
||||
for word in ["first", "second", "third", "fourth"]:
|
||||
assert word in combined
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_duplicate_chunks(self):
|
||||
"""Test that chunker does not produce duplicate content."""
|
||||
chunker = RecursiveCharacterChunker(
|
||||
chunk_size=10,
|
||||
chunk_overlap=0,
|
||||
separators=[""],
|
||||
)
|
||||
text = "abcdefghij"
|
||||
result = await chunker.chunk(text)
|
||||
combined = "".join(result)
|
||||
# Combined length should equal original length (no overlap)
|
||||
assert len(combined) == len(text)
|
||||
@@ -0,0 +1,331 @@
|
||||
"""Unit tests for astrbot.core.tools.registry: builtin_tool decorator, tool registration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrbot.core.agent.tool import FunctionTool
|
||||
from astrbot.core.tools.registry import (
|
||||
BuiltinToolConfigCondition,
|
||||
BuiltinToolConfigRule,
|
||||
_BUILTIN_TOOL_CONFIG_RULES,
|
||||
_builtin_tool_classes_by_name,
|
||||
_builtin_tool_names_by_class,
|
||||
_get_config_value,
|
||||
_json_safe,
|
||||
_MISSING,
|
||||
_resolve_builtin_tool_name,
|
||||
builtin_tool,
|
||||
ensure_builtin_tools_loaded,
|
||||
get_builtin_tool_class,
|
||||
get_builtin_tool_config_rule,
|
||||
get_builtin_tool_config_statuses,
|
||||
get_builtin_tool_config_tags,
|
||||
get_builtin_tool_name,
|
||||
iter_builtin_tool_classes,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_registry():
|
||||
"""Clean the global builtin tool registry before and after each test."""
|
||||
before_classes = dict(_builtin_tool_classes_by_name)
|
||||
before_names = dict(_builtin_tool_names_by_class)
|
||||
before_rules = dict(_BUILTIN_TOOL_CONFIG_RULES)
|
||||
yield
|
||||
_builtin_tool_classes_by_name.clear()
|
||||
_builtin_tool_classes_by_name.update(before_classes)
|
||||
_builtin_tool_names_by_class.clear()
|
||||
_builtin_tool_names_by_class.update(before_names)
|
||||
_BUILTIN_TOOL_CONFIG_RULES.clear()
|
||||
_BUILTIN_TOOL_CONFIG_RULES.update(before_rules)
|
||||
|
||||
|
||||
class TestBuiltinToolDecorator:
|
||||
"""builtin_tool decorator registration."""
|
||||
|
||||
def test_decorator_without_call_registers(self):
|
||||
"""Using @builtin_tool without parentheses registers the class."""
|
||||
|
||||
@builtin_tool
|
||||
class MyTool(FunctionTool):
|
||||
name: str = "my_tool"
|
||||
description: str = "My custom tool"
|
||||
|
||||
assert get_builtin_tool_class("my_tool") is MyTool
|
||||
assert get_builtin_tool_name(MyTool) == "my_tool"
|
||||
|
||||
def test_decorator_with_call_registers(self):
|
||||
"""Using @builtin_tool() with parentheses registers the class."""
|
||||
|
||||
@builtin_tool()
|
||||
class AnotherTool(FunctionTool):
|
||||
name: str = "another_tool"
|
||||
description: str = "Another tool"
|
||||
|
||||
assert get_builtin_tool_class("another_tool") is AnotherTool
|
||||
|
||||
def test_decorator_with_config(self):
|
||||
"""@builtin_tool(config={...}) registers with config rules."""
|
||||
|
||||
@builtin_tool(config={"feature.enabled": True, "mode": ("auto", "manual")})
|
||||
class ConfigTool(FunctionTool):
|
||||
name: str = "config_tool"
|
||||
description: str = "Tool with config"
|
||||
|
||||
rule = get_builtin_tool_config_rule("config_tool")
|
||||
assert rule is not None
|
||||
assert len(rule.conditions) == 2
|
||||
|
||||
def test_name_conflict_raises(self):
|
||||
"""Registering the same name with a different class raises ValueError."""
|
||||
|
||||
@builtin_tool
|
||||
class First(FunctionTool):
|
||||
name: str = "conflict_tool"
|
||||
description: str = "first"
|
||||
|
||||
with pytest.raises(ValueError, match="name conflict"):
|
||||
|
||||
@builtin_tool
|
||||
class Second(FunctionTool):
|
||||
name: str = "conflict_tool"
|
||||
description: str = "second"
|
||||
|
||||
def test_same_class_no_conflict(self):
|
||||
"""Registering the same class again with the same name does not raise."""
|
||||
|
||||
@builtin_tool
|
||||
class SameTool(FunctionTool):
|
||||
name: str = "same_tool"
|
||||
description: str = "same"
|
||||
|
||||
# Registering the same class again should not raise
|
||||
builtin_tool(SameTool)
|
||||
assert get_builtin_tool_class("same_tool") is SameTool
|
||||
|
||||
def test_resolve_tool_name_from_field(self):
|
||||
"""_resolve_builtin_tool_name reads from the 'name' dataclass field."""
|
||||
|
||||
@dataclass
|
||||
class ResolveMe:
|
||||
name: str = "resolved_name"
|
||||
|
||||
# Since ResolveMe is not a FunctionTool, we need to test _resolve_builtin_tool_name
|
||||
# by temporarily making it look like a FunctionTool subclass
|
||||
name = _resolve_builtin_tool_name(ResolveMe)
|
||||
assert name == "resolved_name"
|
||||
|
||||
def test_resolve_tool_name_raises_when_missing(self):
|
||||
"""_resolve_builtin_tool_name raises ValueError when no name is found."""
|
||||
|
||||
class NoName:
|
||||
pass
|
||||
|
||||
with pytest.raises(ValueError, match="does not define a valid name"):
|
||||
_resolve_builtin_tool_name(NoName)
|
||||
|
||||
|
||||
class TestGetAndIter:
|
||||
"""Query functions for the builtin tool registry."""
|
||||
|
||||
def test_get_nonexistent_class_returns_none(self):
|
||||
"""get_builtin_tool_class returns None for unknown names."""
|
||||
assert get_builtin_tool_class("nonexistent_tool") is None
|
||||
|
||||
def test_get_nonexistent_name_returns_none(self):
|
||||
"""get_builtin_tool_name returns None for unknown classes."""
|
||||
class Random(FunctionTool):
|
||||
name: str = "random"
|
||||
description: str = "r"
|
||||
assert get_builtin_tool_name(Random) is None
|
||||
|
||||
def test_iter_builtin_tool_classes_empty(self):
|
||||
"""iter_builtin_tool_classes returns empty tuple when nothing registered."""
|
||||
classes = iter_builtin_tool_classes()
|
||||
# Only pre-existing builtins may be present, but the fixture resets to original state.
|
||||
# We just check it's a tuple.
|
||||
assert isinstance(classes, tuple)
|
||||
|
||||
def test_iter_after_registration(self):
|
||||
"""iter_builtin_tool_classes includes newly registered tools."""
|
||||
|
||||
@builtin_tool
|
||||
class IterTool(FunctionTool):
|
||||
name: str = "iter_tool"
|
||||
description: str = "iter"
|
||||
|
||||
classes = iter_builtin_tool_classes()
|
||||
assert IterTool in classes
|
||||
|
||||
|
||||
class TestConfigCondition:
|
||||
"""BuiltinToolConfigCondition evaluation."""
|
||||
|
||||
def test_equals_condition_match(self):
|
||||
"""'equals' operator returns matched=True when values match."""
|
||||
cond = BuiltinToolConfigCondition(key="enabled", operator="equals", expected=True)
|
||||
result = cond.evaluate({"enabled": True})
|
||||
assert result["matched"] is True
|
||||
|
||||
def test_equals_condition_mismatch(self):
|
||||
"""'equals' operator returns matched=False when values differ."""
|
||||
cond = BuiltinToolConfigCondition(key="enabled", operator="equals", expected=True)
|
||||
result = cond.evaluate({"enabled": False})
|
||||
assert result["matched"] is False
|
||||
|
||||
def test_in_condition_match(self):
|
||||
"""'in' operator returns matched=True when value is in expected."""
|
||||
cond = BuiltinToolConfigCondition(key="mode", operator="in", expected=("a", "b", "c"))
|
||||
result = cond.evaluate({"mode": "b"})
|
||||
assert result["matched"] is True
|
||||
|
||||
def test_in_condition_mismatch(self):
|
||||
"""'in' operator returns matched=False when value is not in expected."""
|
||||
cond = BuiltinToolConfigCondition(key="mode", operator="in", expected=("a", "b"))
|
||||
result = cond.evaluate({"mode": "c"})
|
||||
assert result["matched"] is False
|
||||
|
||||
def test_truthy_condition_match(self):
|
||||
"""'truthy' operator returns matched=True for truthy values."""
|
||||
cond = BuiltinToolConfigCondition(key="timeout", operator="truthy")
|
||||
result = cond.evaluate({"timeout": 30})
|
||||
assert result["matched"] is True
|
||||
|
||||
def test_truthy_condition_mismatch(self):
|
||||
"""'truthy' operator returns matched=False for falsy values."""
|
||||
cond = BuiltinToolConfigCondition(key="timeout", operator="truthy")
|
||||
result = cond.evaluate({"timeout": 0})
|
||||
assert result["matched"] is False
|
||||
|
||||
def test_custom_condition(self):
|
||||
"""'custom' operator delegates to the expected field."""
|
||||
cond = BuiltinToolConfigCondition(key="custom_key", operator="custom", expected=True)
|
||||
result = cond.evaluate({})
|
||||
assert result["matched"] is True
|
||||
|
||||
def test_unsupported_operator_raises(self):
|
||||
"""An unknown operator raises ValueError."""
|
||||
cond = BuiltinToolConfigCondition(key="k", operator="bad_op")
|
||||
with pytest.raises(ValueError, match="Unsupported builtin tool config operator"):
|
||||
cond.evaluate({})
|
||||
|
||||
def test_missing_key_returns_missing(self):
|
||||
"""A key that is not present in config returns _MISSING as actual."""
|
||||
cond = BuiltinToolConfigCondition(key="missing.key", operator="truthy")
|
||||
result = cond.evaluate({})
|
||||
assert result["actual"] is None
|
||||
|
||||
def test_nested_key_access(self):
|
||||
"""_get_config_value traverses dot-separated keys."""
|
||||
config = {"a": {"b": {"c": 42}}}
|
||||
assert _get_config_value(config, "a.b.c") == 42
|
||||
assert _get_config_value(config, "a.b.missing") is _MISSING
|
||||
assert _get_config_value(config, "x") is _MISSING
|
||||
|
||||
def test_json_safe_converts_tuple_to_list(self):
|
||||
"""_json_safe converts tuples to lists."""
|
||||
result = _json_safe((1, (2, 3)))
|
||||
assert result == [1, [2, 3]]
|
||||
|
||||
def test_json_safe_dict(self):
|
||||
"""_json_safe processes dicts recursively."""
|
||||
result = _json_safe({"a": (1, 2), "b": "hello"})
|
||||
assert result == {"a": [1, 2], "b": "hello"}
|
||||
|
||||
|
||||
class TestBuiltinToolConfigRule:
|
||||
"""BuiltinToolConfigRule evaluation."""
|
||||
|
||||
def test_rule_with_conditions(self):
|
||||
"""A rule with conditions evaluates all of them."""
|
||||
c1 = BuiltinToolConfigCondition(key="x", operator="equals", expected=1)
|
||||
c2 = BuiltinToolConfigCondition(key="y", operator="truthy")
|
||||
rule = BuiltinToolConfigRule(conditions=(c1, c2))
|
||||
results = rule.evaluate({"x": 1, "y": True})
|
||||
assert len(results) == 2
|
||||
assert all(r["matched"] for r in results)
|
||||
|
||||
def test_rule_with_evaluator(self):
|
||||
"""A rule with an evaluator callable uses it instead of conditions."""
|
||||
def my_evaluator(config):
|
||||
return [{"key": "custom", "matched": True}]
|
||||
rule = BuiltinToolConfigRule(evaluator=my_evaluator)
|
||||
results = rule.evaluate({"anything": 1})
|
||||
assert results == [{"key": "custom", "matched": True}]
|
||||
|
||||
def test_rule_conditions_are_frozen(self):
|
||||
"""BuiltinToolConfigRule and its conditions are frozen dataclasses."""
|
||||
rule = BuiltinToolConfigRule(conditions=())
|
||||
with pytest.raises(Exception):
|
||||
rule.conditions = ("cannot", "change")
|
||||
|
||||
|
||||
class TestGetBuiltinToolConfigStatuses:
|
||||
"""get_builtin_tool_config_statuses integration."""
|
||||
|
||||
def test_no_rule_returns_empty(self):
|
||||
"""Getting statuses for a tool with no config rule returns []."""
|
||||
statuses = get_builtin_tool_config_statuses("nonexistent", [{"config": {}}])
|
||||
assert statuses == []
|
||||
|
||||
def test_statuses_with_matching_config(self):
|
||||
"""Statuses are returned with enabled=True when all conditions match."""
|
||||
|
||||
@builtin_tool(config={"enabled": True})
|
||||
class StatusTool(FunctionTool):
|
||||
name: str = "status_tool"
|
||||
description: str = "test"
|
||||
|
||||
entries = [{"conf_id": "1", "conf_name": "cfg1", "config": {"enabled": True}}]
|
||||
statuses = get_builtin_tool_config_statuses("status_tool", entries)
|
||||
assert len(statuses) == 1
|
||||
assert statuses[0]["enabled"] is True
|
||||
|
||||
def test_statuses_with_non_matching_config(self):
|
||||
"""Statuses are returned with enabled=False when some conditions fail."""
|
||||
|
||||
@builtin_tool(config={"enabled": True})
|
||||
class StatusTool2(FunctionTool):
|
||||
name: str = "status_tool2"
|
||||
description: str = "test"
|
||||
|
||||
entries = [{"conf_id": "2", "conf_name": "cfg2", "config": {"enabled": False}}]
|
||||
statuses = get_builtin_tool_config_statuses("status_tool2", entries)
|
||||
assert len(statuses) == 1
|
||||
assert statuses[0]["enabled"] is False
|
||||
assert len(statuses[0]["failed_conditions"]) > 0
|
||||
|
||||
def test_get_tags_filters_enabled(self):
|
||||
"""get_builtin_tool_config_tags only returns entries where enabled is True."""
|
||||
|
||||
@builtin_tool(config={"enabled": True})
|
||||
class TagTool(FunctionTool):
|
||||
name: str = "tag_tool"
|
||||
description: str = "test"
|
||||
|
||||
entries = [
|
||||
{"conf_id": "1", "conf_name": "on", "config": {"enabled": True}},
|
||||
{"conf_id": "2", "conf_name": "off", "config": {"enabled": False}},
|
||||
]
|
||||
tags = get_builtin_tool_config_tags("tag_tool", entries)
|
||||
assert len(tags) == 1
|
||||
assert tags[0]["conf_id"] == "1"
|
||||
|
||||
|
||||
class TestEnsureBuiltinToolsLoaded:
|
||||
"""ensure_builtin_tools_loaded idempotency."""
|
||||
|
||||
def test_load_is_idempotent(self):
|
||||
"""Calling ensure_builtin_tools_loaded twice does not raise."""
|
||||
# The first call may fail if the builtin modules have missing deps; catch.
|
||||
try:
|
||||
ensure_builtin_tools_loaded()
|
||||
except Exception:
|
||||
pass
|
||||
# Second call should also not raise (just returns early).
|
||||
ensure_builtin_tools_loaded()
|
||||
@@ -0,0 +1,91 @@
|
||||
"""Unit tests for astrbot.core.agent.run_context: ContextWrapper."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.agent.run_context import ContextWrapper, NoContext
|
||||
from astrbot.core.agent.message import Message
|
||||
|
||||
|
||||
class TestContextWrapper:
|
||||
"""ContextWrapper construction and default values."""
|
||||
|
||||
def test_default_messages_is_empty_list(self):
|
||||
"""ContextWrapper starts with an empty messages list."""
|
||||
ctx = ContextWrapper(context="test")
|
||||
assert ctx.messages == []
|
||||
|
||||
def test_default_tool_call_timeout(self):
|
||||
"""Default tool_call_timeout is 120 seconds."""
|
||||
ctx = ContextWrapper(context="test")
|
||||
assert ctx.tool_call_timeout == 120
|
||||
|
||||
def test_context_string(self):
|
||||
"""context holds a string value."""
|
||||
ctx = ContextWrapper(context="hello_world")
|
||||
assert ctx.context == "hello_world"
|
||||
|
||||
def test_context_integer(self):
|
||||
"""context holds an integer value."""
|
||||
ctx = ContextWrapper(context=42)
|
||||
assert ctx.context == 42
|
||||
|
||||
def test_context_dict(self):
|
||||
"""context holds a dict value."""
|
||||
data = {"key": "value", "num": 1}
|
||||
ctx = ContextWrapper(context=data)
|
||||
assert ctx.context == data
|
||||
|
||||
def test_context_none(self):
|
||||
"""context holds None."""
|
||||
ctx = ContextWrapper(context=None)
|
||||
assert ctx.context is None
|
||||
|
||||
def test_messages_can_be_appended(self):
|
||||
"""messages list supports appending Message objects."""
|
||||
ctx = ContextWrapper(context="test")
|
||||
msg = Message(role="user", content="hi")
|
||||
ctx.messages.append(msg)
|
||||
assert len(ctx.messages) == 1
|
||||
assert ctx.messages[0].content == "hi"
|
||||
|
||||
def test_messages_can_be_replaced(self):
|
||||
"""messages field can be replaced with a new list."""
|
||||
msgs = [Message(role="user", content="a"), Message(role="assistant", content="b")]
|
||||
ctx = ContextWrapper(context="test", messages=msgs)
|
||||
assert len(ctx.messages) == 2
|
||||
|
||||
def test_custom_tool_call_timeout(self):
|
||||
"""tool_call_timeout can be customized."""
|
||||
ctx = ContextWrapper(context="test", tool_call_timeout=30)
|
||||
assert ctx.tool_call_timeout == 30
|
||||
|
||||
def test_generic_type_parameter(self):
|
||||
"""ContextWrapper is generic; it accepts a type parameter."""
|
||||
ctx: ContextWrapper[str] = ContextWrapper(context="typed")
|
||||
assert ctx.context == "typed"
|
||||
|
||||
|
||||
class TestNoContext:
|
||||
"""NoContext is a ContextWrapper[None] singleton-like pattern."""
|
||||
|
||||
def test_no_context_is_contextwrapper(self):
|
||||
"""NoContext is an instance of ContextWrapper."""
|
||||
nc = NoContext()
|
||||
assert isinstance(nc, ContextWrapper)
|
||||
|
||||
def test_no_context_context_is_none(self):
|
||||
"""NoContext always carries context=None."""
|
||||
nc = NoContext()
|
||||
assert nc.context is None
|
||||
|
||||
def test_no_context_messages_empty(self):
|
||||
"""NoContext starts with empty messages."""
|
||||
nc = NoContext()
|
||||
assert nc.messages == []
|
||||
|
||||
def test_no_context_default_timeout(self):
|
||||
"""NoContext has the default 120s timeout."""
|
||||
nc = NoContext()
|
||||
assert nc.tool_call_timeout == 120
|
||||
@@ -0,0 +1,829 @@
|
||||
"""
|
||||
Unit tests for SkillManager.
|
||||
|
||||
Covers construction, list_skills with various combinations (active_only,
|
||||
runtime modes, sandbox caching), set_sandbox_skills_cache,
|
||||
get_sandbox_skills_cache_status, is_sandbox_only_skill, set_skill_active,
|
||||
delete_skill, and config persistence.
|
||||
All tests use mocks to isolate SkillManager from the filesystem and I/O.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.skills.skill_manager import (
|
||||
DEFAULT_SKILLS_CONFIG,
|
||||
SANDBOX_SKILLS_CACHE_FILENAME,
|
||||
SKILLS_CONFIG_FILENAME,
|
||||
SANDBOX_SKILLS_ROOT,
|
||||
SANDBOX_WORKSPACE_ROOT,
|
||||
SkillInfo,
|
||||
SkillManager,
|
||||
_normalize_skill_name,
|
||||
_normalize_cached_sandbox_skill_path,
|
||||
_is_ignored_zip_entry,
|
||||
_sanitize_prompt_path_for_prompt,
|
||||
_sanitize_prompt_description,
|
||||
_sanitize_skill_display_name,
|
||||
build_skills_prompt,
|
||||
_parse_frontmatter,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_astrbot_paths():
|
||||
"""Create mock AstrbotPaths with in-memory-like attributes."""
|
||||
paths = MagicMock()
|
||||
paths.skills = "/tmp/.astrbot/skills"
|
||||
paths.config = Path("/tmp/.astrbot/config")
|
||||
paths.data = Path("/tmp/.astrbot/data")
|
||||
paths.temp = Path("/tmp/.astrbot/temp")
|
||||
return paths
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def skill_manager(mock_astrbot_paths):
|
||||
"""Create a SkillManager with mocked paths and no real FS side effects."""
|
||||
with (
|
||||
patch("os.makedirs") as mock_makedirs,
|
||||
patch.object(Path, "iterdir", return_value=[]),
|
||||
):
|
||||
mgr = SkillManager(
|
||||
skills_root="/tmp/test_skills",
|
||||
astrbot_paths=mock_astrbot_paths,
|
||||
)
|
||||
return mgr
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sandbox_skill_entry():
|
||||
"""Create a canned sandbox skill entry as found in cache."""
|
||||
return {
|
||||
"name": "sandbox-skill",
|
||||
"description": "A sandbox preset skill",
|
||||
"path": f"{SANDBOX_WORKSPACE_ROOT}/{SANDBOX_SKILLS_ROOT}/sandbox-skill/SKILL.md",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# Construction
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillManagerConstruction:
|
||||
"""Test SkillManager construction."""
|
||||
|
||||
def test_construction_sets_attributes(self, mock_astrbot_paths):
|
||||
"""Test that construction sets expected attributes and creates directories."""
|
||||
with patch("os.makedirs") as mock_makedirs:
|
||||
mgr = SkillManager(
|
||||
skills_root="/tmp/test_skills",
|
||||
astrbot_paths=mock_astrbot_paths,
|
||||
)
|
||||
assert mgr.skills_root == "/tmp/test_skills"
|
||||
assert mgr.astrbot_paths is mock_astrbot_paths
|
||||
mock_makedirs.assert_called_once_with("/tmp/test_skills", exist_ok=True)
|
||||
|
||||
def test_construction_default_skills_root(self, mock_astrbot_paths):
|
||||
"""Test that skills_root defaults to astrbot_paths.skills."""
|
||||
with patch("os.makedirs"):
|
||||
mgr = SkillManager(astrbot_paths=mock_astrbot_paths)
|
||||
assert mgr.skills_root == str(mock_astrbot_paths.skills)
|
||||
|
||||
def test_construction_config_path(self, mock_astrbot_paths):
|
||||
"""Test that config_path is derived from astrbot_paths."""
|
||||
with patch("os.makedirs"):
|
||||
mgr = SkillManager(
|
||||
skills_root="/tmp/foo",
|
||||
astrbot_paths=mock_astrbot_paths,
|
||||
)
|
||||
expected = str(mock_astrbot_paths.config / SKILLS_CONFIG_FILENAME)
|
||||
assert mgr.config_path == expected
|
||||
|
||||
def test_construction_sandbox_cache_path(self, mock_astrbot_paths):
|
||||
"""Test that sandbox_skills_cache_path is derived correctly."""
|
||||
with patch("os.makedirs"):
|
||||
mgr = SkillManager(
|
||||
skills_root="/tmp/foo",
|
||||
astrbot_paths=mock_astrbot_paths,
|
||||
)
|
||||
expected = str(
|
||||
mock_astrbot_paths.data / SANDBOX_SKILLS_CACHE_FILENAME,
|
||||
)
|
||||
assert mgr.sandbox_skills_cache_path == expected
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# list_skills - config loading
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillManagerListSkillsConfig:
|
||||
"""Test list_skills config loading behavior."""
|
||||
|
||||
def test_list_skills_creates_default_config_when_missing(
|
||||
self,
|
||||
skill_manager,
|
||||
mock_astrbot_paths,
|
||||
):
|
||||
"""Test that list_skills creates a default config when config file is missing."""
|
||||
# Simulate config file does not exist
|
||||
with (
|
||||
patch("os.path.exists", return_value=False) as mock_exists,
|
||||
patch("builtins.open", mock_open()) as mock_file,
|
||||
patch.object(Path, "iterdir", return_value=[]),
|
||||
patch.object(Path, "exists", return_value=True),
|
||||
):
|
||||
# exists is called multiple times; handle skills_root check too
|
||||
mock_exists.side_effect = lambda p: False # all False
|
||||
|
||||
with patch.object(Path, "mkdir"):
|
||||
result = skill_manager.list_skills()
|
||||
|
||||
assert result == []
|
||||
# Default config should have been saved
|
||||
mock_file.assert_any_call(
|
||||
skill_manager.config_path,
|
||||
"w",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
def test_list_skills_loads_existing_skills_from_config(
|
||||
self,
|
||||
skill_manager,
|
||||
mock_astrbot_paths,
|
||||
):
|
||||
"""Test that list_skills reads skills from skill config files on disk."""
|
||||
config_data = json.dumps({"skills": {"my-skill": {"active": True}}})
|
||||
|
||||
skill_dir = MagicMock(spec=Path)
|
||||
skill_dir.is_dir.return_value = True
|
||||
skill_dir.name = "my-skill"
|
||||
|
||||
skill_md = MagicMock(spec=Path)
|
||||
skill_md.exists.return_value = True
|
||||
skill_md.read_text.return_value = (
|
||||
"---\nname: my-skill\ndescription: My test skill\n---\nDo stuff."
|
||||
)
|
||||
|
||||
with (
|
||||
patch("os.path.exists", return_value=True),
|
||||
patch("builtins.open", mock_open(read_data=config_data)),
|
||||
patch.object(Path, "iterdir", return_value=[skill_dir]),
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
return_value=skill_md,
|
||||
),
|
||||
):
|
||||
result = skill_manager.list_skills()
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].name == "my-skill"
|
||||
assert result[0].description == "My test skill"
|
||||
assert result[0].active is True
|
||||
assert result[0].local_exists is True
|
||||
assert result[0].source_type == "local_only"
|
||||
|
||||
def test_list_skills_skips_directories_without_skill_md(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that directories without SKILL.md are skipped."""
|
||||
config_data = json.dumps({"skills": {}})
|
||||
skill_dir = MagicMock(spec=Path)
|
||||
skill_dir.is_dir.return_value = True
|
||||
skill_dir.name = "not-a-skill"
|
||||
|
||||
with (
|
||||
patch("os.path.exists", return_value=True),
|
||||
patch("builtins.open", mock_open(read_data=config_data)),
|
||||
patch.object(Path, "iterdir", return_value=[skill_dir]),
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
return_value=None,
|
||||
),
|
||||
):
|
||||
result = skill_manager.list_skills()
|
||||
|
||||
assert result == []
|
||||
|
||||
def test_list_skills_skips_non_directory_entries(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that non-directory entries in skills_root are skipped."""
|
||||
config_data = json.dumps({"skills": {}})
|
||||
file_entry = MagicMock(spec=Path)
|
||||
file_entry.is_dir.return_value = False
|
||||
file_entry.name = "README.md"
|
||||
|
||||
with (
|
||||
patch("os.path.exists", return_value=True),
|
||||
patch("builtins.open", mock_open(read_data=config_data)),
|
||||
patch.object(Path, "iterdir", return_value=[file_entry]),
|
||||
):
|
||||
result = skill_manager.list_skills()
|
||||
|
||||
assert result == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# list_skills - active_only filter
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillManagerActiveOnly:
|
||||
"""Test the active_only filter in list_skills."""
|
||||
|
||||
def test_list_skills_active_only_excludes_inactive(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that active_only=True excludes inactive skills."""
|
||||
config_data = json.dumps(
|
||||
{
|
||||
"skills": {
|
||||
"active-skill": {"active": True},
|
||||
"inactive-skill": {"active": False},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
active_md = MagicMock(spec=Path)
|
||||
active_md.read_text.return_value = "---\nname: active-skill\ndescription: Active\n---"
|
||||
inactive_md = MagicMock(spec=Path)
|
||||
inactive_md.read_text.return_value = (
|
||||
"---\nname: inactive-skill\ndescription: Inactive\n---"
|
||||
)
|
||||
|
||||
def normalize_skill_markdown_path(skill_dir):
|
||||
if skill_dir.name == "active-skill":
|
||||
return active_md
|
||||
if skill_dir.name == "inactive-skill":
|
||||
return inactive_md
|
||||
return None
|
||||
|
||||
active_dir = MagicMock(spec=Path)
|
||||
active_dir.is_dir.return_value = True
|
||||
active_dir.name = "active-skill"
|
||||
|
||||
inactive_dir = MagicMock(spec=Path)
|
||||
inactive_dir.is_dir.return_value = True
|
||||
inactive_dir.name = "inactive-skill"
|
||||
|
||||
with (
|
||||
patch("os.path.exists", return_value=True),
|
||||
patch("builtins.open", mock_open(read_data=config_data)),
|
||||
patch.object(
|
||||
Path,
|
||||
"iterdir",
|
||||
return_value=[active_dir, inactive_dir],
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
side_effect=normalize_skill_markdown_path,
|
||||
),
|
||||
):
|
||||
result = skill_manager.list_skills(active_only=True)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].name == "active-skill"
|
||||
|
||||
def test_list_skills_active_only_false_returns_all(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that active_only=False returns all skills."""
|
||||
config_data = json.dumps(
|
||||
{
|
||||
"skills": {
|
||||
"active-skill": {"active": True},
|
||||
"inactive-skill": {"active": False},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
def make_skill_dir(name, md):
|
||||
d = MagicMock(spec=Path)
|
||||
d.is_dir.return_value = True
|
||||
d.name = name
|
||||
return d
|
||||
|
||||
def normalize_skill_markdown_path(skill_dir):
|
||||
md = MagicMock(spec=Path)
|
||||
md.read_text.return_value = f"---\nname: {skill_dir.name}\ndescription: desc\n---"
|
||||
return md
|
||||
|
||||
with (
|
||||
patch("os.path.exists", return_value=True),
|
||||
patch("builtins.open", mock_open(read_data=config_data)),
|
||||
patch.object(
|
||||
Path,
|
||||
"iterdir",
|
||||
return_value=[
|
||||
make_skill_dir("active-skill", None),
|
||||
make_skill_dir("inactive-skill", None),
|
||||
],
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
side_effect=normalize_skill_markdown_path,
|
||||
),
|
||||
):
|
||||
result = skill_manager.list_skills(active_only=False)
|
||||
|
||||
assert len(result) == 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# list_skills - sandbox runtime
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillManagerSandbox:
|
||||
"""Test list_skills with runtime='sandbox'."""
|
||||
|
||||
def test_list_skills_sandbox_includes_sandbox_only_skills(
|
||||
self,
|
||||
skill_manager,
|
||||
sandbox_skill_entry,
|
||||
):
|
||||
"""Test that sandbox runtime includes sandbox-only skills from cache."""
|
||||
config_data = json.dumps({"skills": {}})
|
||||
|
||||
cache_data = json.dumps(
|
||||
{
|
||||
"version": 1,
|
||||
"skills": [sandbox_skill_entry],
|
||||
},
|
||||
)
|
||||
|
||||
def exists_side_effect(path):
|
||||
"""Config file exists, sandbox cache exists."""
|
||||
return True
|
||||
|
||||
def open_side_effect(path, *args, **kwargs):
|
||||
if SANDBOX_SKILLS_CACHE_FILENAME in path:
|
||||
return mock_open(read_data=cache_data).return_value
|
||||
return mock_open(read_data=config_data).return_value
|
||||
|
||||
with (
|
||||
patch("os.path.exists", side_effect=exists_side_effect),
|
||||
patch("builtins.open", side_effect=open_side_effect),
|
||||
patch.object(Path, "iterdir", return_value=[]),
|
||||
):
|
||||
result = skill_manager.list_skills(runtime="sandbox")
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].name == "sandbox-skill"
|
||||
assert result[0].source_type == "sandbox_only"
|
||||
assert result[0].local_exists is False
|
||||
assert result[0].sandbox_exists is True
|
||||
|
||||
def test_list_skills_sandbox_marks_both_synced(
|
||||
self,
|
||||
skill_manager,
|
||||
sandbox_skill_entry,
|
||||
):
|
||||
"""Test that skills existing both locally and in sandbox cache are 'both'."""
|
||||
config_data = json.dumps(
|
||||
{"skills": {"local-skill": {"active": True}}},
|
||||
)
|
||||
|
||||
cache_data = json.dumps(
|
||||
{
|
||||
"version": 1,
|
||||
"skills": [
|
||||
{
|
||||
"name": "local-skill",
|
||||
"description": "Synced skill",
|
||||
"path": f"{SANDBOX_WORKSPACE_ROOT}/{SANDBOX_SKILLS_ROOT}/local-skill/SKILL.md",
|
||||
},
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
skill_dir = MagicMock(spec=Path)
|
||||
skill_dir.is_dir.return_value = True
|
||||
skill_dir.name = "local-skill"
|
||||
|
||||
skill_md = MagicMock(spec=Path)
|
||||
skill_md.read_text.return_value = (
|
||||
"---\nname: local-skill\ndescription: Local desc\n---"
|
||||
)
|
||||
|
||||
with (
|
||||
patch("os.path.exists", return_value=True),
|
||||
patch("builtins.open") as mock_open_func,
|
||||
patch.object(Path, "iterdir", return_value=[skill_dir]),
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
return_value=skill_md,
|
||||
),
|
||||
):
|
||||
# Return different content based on file path
|
||||
def open_side_effect(path, *args, **kwargs):
|
||||
if SANDBOX_SKILLS_CACHE_FILENAME in str(path):
|
||||
return mock_open(read_data=cache_data).return_value
|
||||
return mock_open(read_data=config_data).return_value
|
||||
|
||||
mock_open_func.side_effect = open_side_effect
|
||||
result = skill_manager.list_skills(runtime="sandbox")
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].name == "local-skill"
|
||||
assert result[0].source_type == "both"
|
||||
assert result[0].local_exists is True
|
||||
assert result[0].sandbox_exists is True
|
||||
|
||||
def test_list_skills_sandbox_excludes_local_when_not_in_cache(
|
||||
self,
|
||||
skill_manager,
|
||||
sandbox_skill_entry,
|
||||
):
|
||||
"""Test that local skills without a sandbox cache entry have sandbox_exists=False."""
|
||||
config_data = json.dumps({"skills": {"local-only": {"active": True}}})
|
||||
cache_data = json.dumps({"version": 1, "skills": []})
|
||||
|
||||
skill_dir = MagicMock(spec=Path)
|
||||
skill_dir.is_dir.return_value = True
|
||||
skill_dir.name = "local-only"
|
||||
|
||||
skill_md = MagicMock(spec=Path)
|
||||
skill_md.read_text.return_value = (
|
||||
"---\nname: local-only\ndescription: Local only\n---"
|
||||
)
|
||||
|
||||
with (
|
||||
patch("os.path.exists", return_value=True),
|
||||
patch("builtins.open") as mock_open_func,
|
||||
patch.object(Path, "iterdir", return_value=[skill_dir]),
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
return_value=skill_md,
|
||||
),
|
||||
):
|
||||
def open_side_effect(path, *args, **kwargs):
|
||||
if SANDBOX_SKILLS_CACHE_FILENAME in str(path):
|
||||
return mock_open(read_data=cache_data).return_value
|
||||
return mock_open(read_data=config_data).return_value
|
||||
|
||||
mock_open_func.side_effect = open_side_effect
|
||||
result = skill_manager.list_skills(runtime="sandbox")
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].name == "local-only"
|
||||
assert result[0].local_exists is True
|
||||
assert result[0].sandbox_exists is False
|
||||
assert result[0].source_type == "local_only"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# set_sandbox_skills_cache and get_sandbox_skills_cache_status
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillManagerSandboxCache:
|
||||
"""Test sandbox cache management."""
|
||||
|
||||
def test_set_sandbox_skills_cache_saves_deduped(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that set_sandbox_skills_cache deduplicates and saves."""
|
||||
skills = [
|
||||
{"name": "skill-a", "description": "A", "path": "/workspace/.../SKILL.md"},
|
||||
{"name": "skill-b", "description": "B", "path": "/workspace/.../SKILL.md"},
|
||||
{"name": "skill-a", "description": "A dup", "path": "/workspace/.../SKILL.md"},
|
||||
]
|
||||
|
||||
with patch.object(skill_manager, "_save_sandbox_skills_cache") as mock_save:
|
||||
skill_manager.set_sandbox_skills_cache(skills)
|
||||
|
||||
saved = mock_save.call_args[0][0]
|
||||
assert "skills" in saved
|
||||
assert len(saved["skills"]) == 2
|
||||
names = [s["name"] for s in saved["skills"]]
|
||||
assert "skill-a" in names
|
||||
assert "skill-b" in names
|
||||
|
||||
def test_set_sandbox_skills_cache_skips_invalid_names(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that entries with invalid names are skipped."""
|
||||
skills = [
|
||||
{"name": "valid-skill", "description": "ok", "path": ""},
|
||||
{"name": "", "description": "empty", "path": ""},
|
||||
{"name": "../escape", "description": "bad", "path": ""},
|
||||
]
|
||||
|
||||
with patch.object(skill_manager, "_save_sandbox_skills_cache") as mock_save:
|
||||
skill_manager.set_sandbox_skills_cache(skills)
|
||||
|
||||
saved = mock_save.call_args[0][0]
|
||||
assert len(saved["skills"]) == 1
|
||||
assert saved["skills"][0]["name"] == "valid-skill"
|
||||
|
||||
def test_get_sandbox_skills_cache_status_ready(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that status returns ready=True when cache has skills."""
|
||||
with (
|
||||
patch.object(
|
||||
skill_manager,
|
||||
"_load_sandbox_skills_cache",
|
||||
return_value={
|
||||
"version": 1,
|
||||
"skills": [{"name": "s1"}],
|
||||
"updated_at": "2025-01-01T00:00:00+00:00",
|
||||
},
|
||||
),
|
||||
patch("os.path.exists", return_value=True),
|
||||
):
|
||||
status = skill_manager.get_sandbox_skills_cache_status()
|
||||
|
||||
assert status["exists"] is True
|
||||
assert status["ready"] is True
|
||||
assert status["count"] == 1
|
||||
|
||||
def test_get_sandbox_skills_cache_status_not_ready(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that status returns ready=False when no skills cached."""
|
||||
with (
|
||||
patch.object(
|
||||
skill_manager,
|
||||
"_load_sandbox_skills_cache",
|
||||
return_value={"version": 1, "skills": []},
|
||||
),
|
||||
patch("os.path.exists", return_value=True),
|
||||
):
|
||||
status = skill_manager.get_sandbox_skills_cache_status()
|
||||
|
||||
assert status["exists"] is True
|
||||
assert status["ready"] is False
|
||||
assert status["count"] == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# is_sandbox_only_skill
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillManagerIsSandboxOnly:
|
||||
"""Test is_sandbox_only_skill."""
|
||||
|
||||
def test_is_sandbox_only_skill_returns_true_when_only_in_cache(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that a skill existing only in cache returns True."""
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
skill_manager,
|
||||
"_load_sandbox_skills_cache",
|
||||
return_value={
|
||||
"version": 1,
|
||||
"skills": [{"name": "sandbox-only"}],
|
||||
},
|
||||
),
|
||||
):
|
||||
assert skill_manager.is_sandbox_only_skill("sandbox-only") is True
|
||||
|
||||
def test_is_sandbox_only_skill_returns_false_when_local_exists(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that a skill with local SKILL.md returns False."""
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
):
|
||||
assert skill_manager.is_sandbox_only_skill("local-skill") is False
|
||||
|
||||
def test_is_sandbox_only_skill_returns_false_for_nonexistent(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that a skill not in cache or local returns False."""
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
skill_manager,
|
||||
"_load_sandbox_skills_cache",
|
||||
return_value={"version": 1, "skills": []},
|
||||
),
|
||||
):
|
||||
assert skill_manager.is_sandbox_only_skill("ghost") is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# set_skill_active / delete_skill
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillManagerMutations:
|
||||
"""Test set_skill_active and delete_skill."""
|
||||
|
||||
def test_set_skill_active_saves_config(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that set_skill_active writes the config with the new active state."""
|
||||
existing_config = {"skills": {"my-skill": {"active": False}}}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.skills.skill_manager._normalize_skill_markdown_path",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch("os.path.exists", return_value=True),
|
||||
patch("builtins.open", mock_open(read_data=json.dumps(existing_config))),
|
||||
):
|
||||
skill_manager.set_skill_active("my-skill", True)
|
||||
|
||||
# Config should have been saved with active=True
|
||||
# We verify by checking that config file was opened for writing
|
||||
# with content containing "active": true
|
||||
# The exact assertion depends on mocking — key is it doesn't crash
|
||||
|
||||
def test_set_skill_active_raises_on_sandbox_only(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that set_skill_active raises PermissionError for sandbox-only skills."""
|
||||
with (
|
||||
patch.object(skill_manager, "is_sandbox_only_skill", return_value=True),
|
||||
pytest.raises(PermissionError, match="Sandbox preset skill"),
|
||||
):
|
||||
skill_manager.set_skill_active("sandbox-only", False)
|
||||
|
||||
def test_delete_skill_removes_directory_and_config(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that delete_skill removes the skill directory and config entry."""
|
||||
from astrbot.core.skills.skill_manager import _normalize_skill_markdown_path
|
||||
|
||||
skill_dir = MagicMock(spec=Path)
|
||||
skill_dir.exists.return_value = True
|
||||
skill_dir.name = "my-skill"
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
skill_manager,
|
||||
"is_sandbox_only_skill",
|
||||
return_value=False,
|
||||
),
|
||||
patch("pathlib.Path", return_value=skill_dir) as mock_path_cls,
|
||||
patch("shutil.rmtree") as mock_rmtree,
|
||||
patch.object(skill_manager, "_remove_skill_from_sandbox_cache"),
|
||||
patch.object(skill_manager, "_load_config",
|
||||
return_value={"skills": {"my-skill": {"active": True}}}),
|
||||
patch.object(skill_manager, "_save_config") as mock_save,
|
||||
):
|
||||
skill_manager.delete_skill("my-skill")
|
||||
mock_rmtree.assert_called_once_with(skill_dir)
|
||||
mock_save.assert_called_once()
|
||||
|
||||
def test_delete_skill_raises_on_sandbox_only(
|
||||
self,
|
||||
skill_manager,
|
||||
):
|
||||
"""Test that delete_skill raises PermissionError for sandbox-only skills."""
|
||||
with (
|
||||
patch.object(skill_manager, "is_sandbox_only_skill", return_value=True),
|
||||
pytest.raises(PermissionError, match="Sandbox preset skill"),
|
||||
):
|
||||
skill_manager.delete_skill("sandbox-only")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# utility / helper functions
|
||||
# ---------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSkillManagerUtilities:
|
||||
"""Test standalone utility functions in skill_manager module."""
|
||||
|
||||
def test_normalize_skill_name_replaces_spaces(self):
|
||||
"""Test that _normalize_skill_name replaces whitespace with underscores."""
|
||||
result = _normalize_skill_name("my cool skill")
|
||||
assert result == "my_cool_skill"
|
||||
assert _normalize_skill_name(" hello world ") == "hello_world"
|
||||
assert _normalize_skill_name(None) == ""
|
||||
|
||||
def test_normalize_cached_sandbox_skill_path_default(self):
|
||||
"""Test that empty path defaults to the standard sandbox SKILL.md path."""
|
||||
result = _normalize_cached_sandbox_skill_path("my-skill", "")
|
||||
expected = f"{SANDBOX_WORKSPACE_ROOT}/{SANDBOX_SKILLS_ROOT}/my-skill/SKILL.md"
|
||||
assert result == expected
|
||||
|
||||
def test_normalize_cached_sandbox_skill_path_rejects_relative_escape(self):
|
||||
"""Test that path with '..' is rejected and falls back to default."""
|
||||
result = _normalize_cached_sandbox_skill_path("my-skill", "/workspace/../../etc/SKILL.md")
|
||||
expected = f"{SANDBOX_WORKSPACE_ROOT}/{SANDBOX_SKILLS_ROOT}/my-skill/SKILL.md"
|
||||
assert result == expected
|
||||
|
||||
def test_normalize_cached_sandbox_skill_path_rejects_wrong_filename(self):
|
||||
"""Test that a path not ending in SKILL.md falls back to default."""
|
||||
result = _normalize_cached_sandbox_skill_path("my-skill", "/workspace/skills/my-skill/README.md")
|
||||
expected = f"{SANDBOX_WORKSPACE_ROOT}/{SANDBOX_SKILLS_ROOT}/my-skill/SKILL.md"
|
||||
assert result == expected
|
||||
|
||||
def test_normalize_cached_sandbox_skill_path_rejects_wrong_dir_name(self):
|
||||
"""Test that a path with mismatched directory name falls back to default."""
|
||||
result = _normalize_cached_sandbox_skill_path("my-skill", "/workspace/skills/other-skill/SKILL.md")
|
||||
expected = f"{SANDBOX_WORKSPACE_ROOT}/{SANDBOX_SKILLS_ROOT}/my-skill/SKILL.md"
|
||||
assert result == expected
|
||||
|
||||
def test_is_ignored_zip_entry_macosx(self):
|
||||
"""Test that __MACOSX entries are ignored."""
|
||||
assert _is_ignored_zip_entry("__MACOSX/") is True
|
||||
assert _is_ignored_zip_entry("__MACOSX/somefile") is True
|
||||
assert _is_ignored_zip_entry("my-skill/SKILL.md") is False
|
||||
|
||||
def test_sanitize_prompt_path_for_prompt_removes_backticks(self):
|
||||
"""Test that backticks are stripped from path."""
|
||||
result = _sanitize_prompt_path_for_prompt("/path/with`backticks`/SKILL.md")
|
||||
assert "`" not in result
|
||||
assert "backticks" in result
|
||||
|
||||
def test_sanitize_prompt_description_cleans_whitespace(self):
|
||||
"""Test that description is cleaned of extra whitespace."""
|
||||
result = _sanitize_prompt_description(" hello world ")
|
||||
assert result == "hello world"
|
||||
|
||||
def test_sanitize_skill_display_name_invalid(self):
|
||||
"""Test that invalid names return <invalid_skill_name>."""
|
||||
result = _sanitize_skill_display_name("../escape")
|
||||
assert result == "<invalid_skill_name>"
|
||||
|
||||
def test_sanitize_skill_display_name_valid(self):
|
||||
"""Test that valid names pass through unchanged."""
|
||||
result = _sanitize_skill_display_name("my-valid_skill1")
|
||||
assert result == "my-valid_skill1"
|
||||
|
||||
def test_parse_frontmatter_extracts_meta(self):
|
||||
"""Test that YAML frontmatter is parsed correctly."""
|
||||
text = "---\nname: test\ndescription: hello\ninput_schema:\n type: object\n---\nBody"
|
||||
result = _parse_frontmatter(text)
|
||||
assert result["name"] == "test"
|
||||
assert result["description"] == "hello"
|
||||
assert result["input_schema"] == {"type": "object"}
|
||||
|
||||
def test_parse_frontmatter_returns_empty_when_no_frontmatter(self):
|
||||
"""Test that text without frontmatter returns empty dict."""
|
||||
result = _parse_frontmatter("No frontmatter here")
|
||||
assert result == {}
|
||||
|
||||
def test_parse_frontmatter_returns_empty_on_yaml_error(self):
|
||||
"""Test that malformed YAML returns empty dict."""
|
||||
text = "---\n: invalid yaml :::\n---\nBody"
|
||||
result = _parse_frontmatter(text)
|
||||
assert result == {}
|
||||
|
||||
def test_build_skills_prompt_generates_block(self):
|
||||
"""Test that build_skills_prompt generates a non-empty prompt block."""
|
||||
skills = [
|
||||
SkillInfo(
|
||||
name="test-skill",
|
||||
description="A test skill",
|
||||
path="/tmp/skills/test-skill/SKILL.md",
|
||||
active=True,
|
||||
),
|
||||
]
|
||||
result = build_skills_prompt(skills)
|
||||
assert "## Skills" in result
|
||||
assert "test-skill" in result
|
||||
assert "A test skill" in result
|
||||
assert "SKILL.md" in result
|
||||
|
||||
def test_build_skills_prompt_empty_list(self):
|
||||
"""Test that build_skills_prompt with no skills still produces a block."""
|
||||
result = build_skills_prompt([])
|
||||
assert "## Skills" in result
|
||||
assert "Available skills" in result
|
||||
@@ -0,0 +1,530 @@
|
||||
"""Unit tests for astrbot.core.star.star_manager.
|
||||
|
||||
Tests PluginManager initialization, load, install, and uninstall flows with
|
||||
full mock isolation (no filesystem, no network).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, PropertyMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.star.star_manager import (
|
||||
PluginDependencyInstallError,
|
||||
PluginManager,
|
||||
PluginVersionIncompatibleError,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PluginVersionIncompatibleError / PluginDependencyInstallError
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPluginExceptions:
|
||||
"""Custom exception classes."""
|
||||
|
||||
def test_version_incompatible_error(self):
|
||||
"""PluginVersionIncompatibleError is a plain exception."""
|
||||
exc = PluginVersionIncompatibleError("bad version")
|
||||
assert isinstance(exc, Exception)
|
||||
assert str(exc) == "bad version"
|
||||
|
||||
def test_dependency_install_error(self):
|
||||
"""PluginDependencyInstallError wraps the original error."""
|
||||
inner = ValueError("pip failed")
|
||||
exc = PluginDependencyInstallError(
|
||||
plugin_label="my_plugin",
|
||||
requirements_path="/path/requirements.txt",
|
||||
error=inner,
|
||||
)
|
||||
assert exc.plugin_label == "my_plugin"
|
||||
assert exc.requirements_path == "/path/requirements.txt"
|
||||
assert exc.error is inner
|
||||
assert "pip failed" in str(exc)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_context():
|
||||
"""A fully mocked Context with required methods."""
|
||||
ctx = MagicMock()
|
||||
ctx.get_all_stars.return_value = []
|
||||
ctx.get_registered_star.return_value = None
|
||||
return ctx
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_config():
|
||||
"""A mock AstrBotConfig."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def plugin_manager(mock_context, mock_config):
|
||||
"""Create a PluginManager with all dependencies mocked."""
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.get_astrbot_plugin_path",
|
||||
return_value="/mock/plugins",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.get_astrbot_config_path",
|
||||
return_value="/mock/config",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.get_astrbot_path",
|
||||
return_value="/mock/astrbot",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.PluginUpdator",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.StarTools",
|
||||
),
|
||||
):
|
||||
pm = PluginManager(mock_context, mock_config)
|
||||
# Disable hot-reload watcher
|
||||
pm.tasks = set()
|
||||
return pm
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constructor / Init
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPluginManagerInit:
|
||||
"""PluginManager constructor behavior."""
|
||||
|
||||
def test_init_sets_paths(self, plugin_manager):
|
||||
"""Constructor sets correct default paths."""
|
||||
assert plugin_manager.plugin_store_path == "/mock/plugins"
|
||||
assert plugin_manager.plugin_config_path == "/mock/config"
|
||||
assert "astrbot" in plugin_manager.reserved_plugin_path
|
||||
|
||||
def test_init_sets_empty_failed_dict(self, plugin_manager):
|
||||
"""Constructor initializes empty failed plugin tracking."""
|
||||
assert plugin_manager.failed_plugin_dict == {}
|
||||
assert plugin_manager.failed_plugin_info == ""
|
||||
|
||||
def test_init_stores_context_and_config(self, plugin_manager, mock_context, mock_config):
|
||||
"""Constructor stores the context and config references."""
|
||||
assert plugin_manager.context is mock_context
|
||||
assert plugin_manager.config is mock_config
|
||||
|
||||
def test_init_sets_lock(self, plugin_manager):
|
||||
"""Constructor creates an asyncio.Lock."""
|
||||
import asyncio
|
||||
|
||||
assert isinstance(plugin_manager._pm_lock, asyncio.Lock)
|
||||
|
||||
def test_init_starts_watcher_when_reload_env_set(self, mock_context, mock_config):
|
||||
"""When ASTRBOT_RELOAD=1, the file watcher task is created."""
|
||||
with (
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.get_astrbot_plugin_path",
|
||||
return_value="/mock/plugins",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.get_astrbot_config_path",
|
||||
return_value="/mock/config",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.get_astrbot_path",
|
||||
return_value="/mock",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.PluginUpdator",
|
||||
),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.StarTools",
|
||||
),
|
||||
patch.dict("os.environ", {"ASTRBOT_RELOAD": "1"}),
|
||||
patch(
|
||||
"astrbot.core.star.star_manager.asyncio.create_task",
|
||||
) as mock_create_task,
|
||||
):
|
||||
pm = PluginManager(mock_context, mock_config)
|
||||
# Clean up after init
|
||||
assert mock_create_task.called
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _get_classes / _get_modules / _get_plugin_modules (static helpers)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetClasses:
|
||||
"""PluginManager._get_classes() static helper."""
|
||||
|
||||
def test_returns_class_names_from_module(self):
|
||||
"""_get_classes finds classes ending in 'plugin' or named 'main'."""
|
||||
import types
|
||||
|
||||
module = types.ModuleType("test_module")
|
||||
exec(
|
||||
"""
|
||||
class MyPlugin:
|
||||
pass
|
||||
|
||||
class NotAPlugin:
|
||||
pass
|
||||
""",
|
||||
module.__dict__,
|
||||
)
|
||||
result = PluginManager._get_classes(module)
|
||||
assert "MyPlugin" in result
|
||||
assert "NotAPlugin" not in result
|
||||
|
||||
def test_returns_main_class(self):
|
||||
"""_get_classes includes a class named 'main' (case-insensitive)."""
|
||||
import types
|
||||
|
||||
module = types.ModuleType("test_module")
|
||||
exec(
|
||||
"""
|
||||
class Main:
|
||||
pass
|
||||
""",
|
||||
module.__dict__,
|
||||
)
|
||||
result = PluginManager._get_classes(module)
|
||||
assert "Main" in result
|
||||
|
||||
def test_returns_empty_when_no_match(self):
|
||||
"""_get_classes returns empty list when no matching classes exist."""
|
||||
import types
|
||||
|
||||
module = types.ModuleType("test_module")
|
||||
exec(
|
||||
"""
|
||||
class Helper:
|
||||
pass
|
||||
""",
|
||||
module.__dict__,
|
||||
)
|
||||
result = PluginManager._get_classes(module)
|
||||
assert result == []
|
||||
|
||||
|
||||
class TestValidateImportableName:
|
||||
"""PluginManager._validate_importable_name() static helper."""
|
||||
|
||||
def test_valid_name_passes(self):
|
||||
"""A valid Python identifier passes validation."""
|
||||
PluginManager._validate_importable_name("my_plugin") # no raise
|
||||
|
||||
def test_rejects_path_separator(self):
|
||||
"""A name containing / raises ValueError."""
|
||||
with pytest.raises(ValueError, match="路径分隔符"):
|
||||
PluginManager._validate_importable_name("my/plugin")
|
||||
|
||||
def test_rejects_invalid_identifier(self):
|
||||
"""A name that is not a valid Python identifier raises Exception."""
|
||||
with pytest.raises(Exception, match="合法的模块名称"):
|
||||
PluginManager._validate_importable_name("123invalid")
|
||||
|
||||
def test_rejects_keyword(self):
|
||||
"""A name that is a Python keyword raises Exception."""
|
||||
with pytest.raises(Exception, match="合法的模块名称"):
|
||||
PluginManager._validate_importable_name("class")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# load
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLoad:
|
||||
"""PluginManager.load() behavior."""
|
||||
|
||||
@patch("astrbot.core.star.star_manager.sp.global_get")
|
||||
@patch("astrbot.core.star.star_manager.sync_command_configs", new_callable=AsyncMock)
|
||||
async def test_load_returns_true_when_no_plugins(
|
||||
self, mock_sync, mock_sp_get, plugin_manager
|
||||
):
|
||||
"""load() returns (True, None) when no plugin modules exist."""
|
||||
mock_sp_get.return_value = []
|
||||
plugin_manager._get_plugin_modules = MagicMock(return_value=[])
|
||||
success, error = await plugin_manager.load()
|
||||
assert success is True
|
||||
assert error is None
|
||||
|
||||
@patch("astrbot.core.star.star_manager.sp.global_get")
|
||||
@patch("astrbot.core.star.star_manager.sync_command_configs", new_callable=AsyncMock)
|
||||
async def test_load_returns_false_when_modules_is_none(
|
||||
self, mock_sync, mock_sp_get, plugin_manager
|
||||
):
|
||||
"""load() returns (False, msg) when _get_plugin_modules returns None."""
|
||||
mock_sp_get.return_value = []
|
||||
plugin_manager._get_plugin_modules = MagicMock(return_value=None)
|
||||
success, error = await plugin_manager.load()
|
||||
assert success is False
|
||||
assert "未找到" in error
|
||||
|
||||
@patch("astrbot.core.star.star_manager.sp.global_get")
|
||||
@patch("astrbot.core.star.star_manager.sync_command_configs", new_callable=AsyncMock)
|
||||
async def test_load_handles_import_failure(
|
||||
self, mock_sync, mock_sp_get, plugin_manager
|
||||
):
|
||||
"""load() records failed plugins when import fails."""
|
||||
mock_sp_get.return_value = []
|
||||
plugin_manager._get_plugin_modules = MagicMock(
|
||||
return_value=[
|
||||
{
|
||||
"pname": "broken_plugin",
|
||||
"module": "main",
|
||||
"module_path": "/mock/plugins/broken_plugin/main",
|
||||
"reserved": False,
|
||||
}
|
||||
]
|
||||
)
|
||||
plugin_manager._import_plugin_with_dependency_recovery = AsyncMock(
|
||||
side_effect=ImportError("Module not found")
|
||||
)
|
||||
plugin_manager._load_plugin_metadata = MagicMock(return_value=None)
|
||||
plugin_manager._build_failed_plugin_record = MagicMock(
|
||||
return_value={"name": "broken_plugin", "error": "Module not found"}
|
||||
)
|
||||
|
||||
success, error = await plugin_manager.load()
|
||||
assert success is False
|
||||
assert "broken_plugin" in plugin_manager.failed_plugin_dict
|
||||
|
||||
@patch("astrbot.core.star.star_manager.sp.global_get")
|
||||
@patch("astrbot.core.star.star_manager.sync_command_configs", new_callable=AsyncMock)
|
||||
async def test_load_calls_sync_command_configs(
|
||||
self, mock_sync, mock_sp_get, plugin_manager
|
||||
):
|
||||
"""load() calls sync_command_configs after processing plugins."""
|
||||
mock_sp_get.return_value = []
|
||||
plugin_manager._get_plugin_modules = MagicMock(return_value=[])
|
||||
await plugin_manager.load()
|
||||
mock_sync.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# install_plugin
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestInstallPlugin:
|
||||
"""PluginManager.install_plugin() behavior."""
|
||||
|
||||
@patch("astrbot.core.star.star_manager.sp.global_get")
|
||||
@patch("astrbot.core.star.star_manager.sync_command_configs", new_callable=AsyncMock)
|
||||
@patch("astrbot.core.star.star_manager.Metric")
|
||||
async def test_install_plugin_success(
|
||||
self, mock_metric, mock_sync, mock_sp_get, plugin_manager, mock_context
|
||||
):
|
||||
"""install_plugin installs and loads a plugin successfully."""
|
||||
mock_sp_get.return_value = []
|
||||
plugin_manager.updator.parse_github_url = MagicMock(
|
||||
return_value=("owner", "test_repo", "repo")
|
||||
)
|
||||
plugin_manager.updator.format_name = MagicMock(return_value="test_repo")
|
||||
|
||||
# Mock the install to return a plugin path
|
||||
plugin_manager.updator.install = AsyncMock(
|
||||
return_value="/mock/plugins/test_repo"
|
||||
)
|
||||
plugin_manager._get_plugin_dir_name_from_metadata = MagicMock(
|
||||
return_value="test_repo"
|
||||
)
|
||||
plugin_manager._ensure_plugin_requirements = AsyncMock()
|
||||
plugin_manager.load = AsyncMock(return_value=(True, None))
|
||||
|
||||
mock_plugin_meta = MagicMock()
|
||||
mock_plugin_meta.repo = "https://github.com/test/test_repo"
|
||||
mock_plugin_meta.name = "test_plugin"
|
||||
mock_context.get_registered_star.side_effect = lambda name: (
|
||||
mock_plugin_meta if name == "test_repo" else None
|
||||
)
|
||||
|
||||
result = await plugin_manager.install_plugin(
|
||||
"https://github.com/test/test_repo"
|
||||
)
|
||||
assert result is not None
|
||||
assert result["repo"] == "https://github.com/test/test_repo"
|
||||
assert result["name"] == "test_plugin"
|
||||
|
||||
@patch("astrbot.core.star.star_manager.sp.global_get")
|
||||
@patch("astrbot.core.star.star_manager.sync_command_configs", new_callable=AsyncMock)
|
||||
@patch("astrbot.core.star.star_manager.Metric")
|
||||
async def test_install_plugin_raises_when_load_fails(
|
||||
self, mock_metric, mock_sync, mock_sp_get, plugin_manager
|
||||
):
|
||||
"""install_plugin raises when load() returns failure."""
|
||||
mock_sp_get.return_value = []
|
||||
plugin_manager.updator.parse_github_url = MagicMock(
|
||||
return_value=("owner", "test_repo", "repo")
|
||||
)
|
||||
plugin_manager.updator.format_name = MagicMock(return_value="test_repo")
|
||||
plugin_manager.updator.install = AsyncMock(
|
||||
return_value="/mock/plugins/test_repo"
|
||||
)
|
||||
plugin_manager._get_plugin_dir_name_from_metadata = MagicMock(
|
||||
return_value="test_repo"
|
||||
)
|
||||
plugin_manager._ensure_plugin_requirements = AsyncMock()
|
||||
plugin_manager.load = AsyncMock(
|
||||
return_value=(False, "Version incompatible error")
|
||||
)
|
||||
|
||||
with pytest.raises(Exception, match="Version incompatible error"):
|
||||
await plugin_manager.install_plugin(
|
||||
"https://github.com/test/test_repo"
|
||||
)
|
||||
|
||||
@patch("astrbot.core.star.star_manager.sp.global_get")
|
||||
@patch("astrbot.core.star.star_manager.sync_command_configs", new_callable=AsyncMock)
|
||||
@patch("astrbot.core.star.star_manager.Metric")
|
||||
async def test_install_plugin_raises_when_dir_exists(
|
||||
self, mock_metric, mock_sync, mock_sp_get, plugin_manager
|
||||
):
|
||||
"""install_plugin raises when the target directory already exists."""
|
||||
mock_sp_get.return_value = []
|
||||
plugin_manager.updator.parse_github_url = MagicMock(
|
||||
return_value=("owner", "test_repo", "repo")
|
||||
)
|
||||
plugin_manager.updator.format_name = MagicMock(return_value="test_repo")
|
||||
|
||||
with (
|
||||
patch("astrbot.core.star.star_manager.anyio.Path") as mock_path,
|
||||
):
|
||||
mock_path_instance = MagicMock()
|
||||
mock_path.return_value = mock_path_instance
|
||||
mock_path_instance.exists = AsyncMock(return_value=True)
|
||||
|
||||
with pytest.raises(Exception, match="已存在"):
|
||||
await plugin_manager.install_plugin(
|
||||
"https://github.com/test/test_repo"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# uninstall_plugin
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestUninstallPlugin:
|
||||
"""PluginManager.uninstall_plugin() behavior."""
|
||||
|
||||
@patch("astrbot.core.star.star_manager.remove_dir")
|
||||
@patch("astrbot.core.star.star_manager.unregister_platform_adapters_by_module")
|
||||
async def test_uninstall_plugin_success(
|
||||
self,
|
||||
mock_unregister,
|
||||
mock_remove_dir,
|
||||
plugin_manager,
|
||||
mock_context,
|
||||
):
|
||||
"""uninstall_plugin terminates, unbinds, and removes plugin directory."""
|
||||
mock_plugin = MagicMock()
|
||||
mock_plugin.reserved = False
|
||||
mock_plugin.root_dir_name = "test_repo"
|
||||
mock_plugin.name = "test_plugin"
|
||||
mock_plugin.module_path = "data.plugins.test_repo.main"
|
||||
mock_context.get_registered_star.return_value = mock_plugin
|
||||
|
||||
plugin_manager._terminate_plugin = AsyncMock()
|
||||
plugin_manager._unbind_plugin = AsyncMock()
|
||||
|
||||
await plugin_manager.uninstall_plugin("test_plugin")
|
||||
|
||||
plugin_manager._terminate_plugin.assert_called_once_with(mock_plugin)
|
||||
plugin_manager._unbind_plugin.assert_called_once_with(
|
||||
"test_plugin", "data.plugins.test_repo.main"
|
||||
)
|
||||
mock_remove_dir.assert_called_once()
|
||||
|
||||
@patch("astrbot.core.star.star_manager.remove_dir")
|
||||
async def test_uninstall_plugin_raises_when_not_found(
|
||||
self, mock_remove_dir, plugin_manager, mock_context
|
||||
):
|
||||
"""uninstall_plugin raises when plugin is not registered."""
|
||||
mock_context.get_registered_star.return_value = None
|
||||
with pytest.raises(Exception, match="插件不存在"):
|
||||
await plugin_manager.uninstall_plugin("nonexistent")
|
||||
|
||||
@patch("astrbot.core.star.star_manager.remove_dir")
|
||||
async def test_uninstall_plugin_raises_for_reserved(
|
||||
self, mock_remove_dir, plugin_manager, mock_context
|
||||
):
|
||||
"""uninstall_plugin raises when plugin is reserved."""
|
||||
mock_plugin = MagicMock()
|
||||
mock_plugin.reserved = True
|
||||
mock_context.get_registered_star.return_value = mock_plugin
|
||||
|
||||
with pytest.raises(Exception, match="保留插件"):
|
||||
await plugin_manager.uninstall_plugin("reserved_plugin")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _validate_astrbot_version_specifier
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestValidateAstrbotVersion:
|
||||
"""PluginManager._validate_astrbot_version_specifier() behavior."""
|
||||
|
||||
def test_none_version_returns_valid(self):
|
||||
"""None version specifier is valid."""
|
||||
valid, msg = PluginManager._validate_astrbot_version_specifier(None)
|
||||
assert valid is True
|
||||
assert msg is None
|
||||
|
||||
@patch("astrbot.core.star.star_manager.VERSION", "4.16.0")
|
||||
def test_version_in_range(self):
|
||||
"""A version specifier that includes the current version is valid."""
|
||||
valid, msg = PluginManager._validate_astrbot_version_specifier(">=4.16,<5")
|
||||
assert valid is True
|
||||
assert msg is None
|
||||
|
||||
@patch("astrbot.core.star.star_manager.VERSION", "4.15.0")
|
||||
def test_version_out_of_range(self):
|
||||
"""A version specifier that excludes the current version is invalid."""
|
||||
valid, msg = PluginManager._validate_astrbot_version_specifier(">=4.16,<5")
|
||||
assert valid is False
|
||||
assert "does not satisfy" in msg
|
||||
|
||||
def test_invalid_specifier_returns_error(self):
|
||||
"""An unparseable specifier returns invalid with error message."""
|
||||
valid, msg = PluginManager._validate_astrbot_version_specifier("not_a_version")
|
||||
assert valid is False
|
||||
assert "PEP 440" in msg
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _load_plugin_metadata
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLoadPluginMetadata:
|
||||
"""PluginManager._load_plugin_metadata() behavior."""
|
||||
|
||||
def test_raises_when_path_does_not_exist(self):
|
||||
"""_load_plugin_metadata raises when plugin_path does not exist."""
|
||||
with patch("astrbot.core.star.star_manager.os.path.exists", return_value=False):
|
||||
with pytest.raises(Exception, match="插件不存在"):
|
||||
PluginManager._load_plugin_metadata("/nonexistent/path")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# reload_failed_plugin
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestReloadFailedPlugin:
|
||||
"""PluginManager.reload_failed_plugin() behavior."""
|
||||
|
||||
async def test_returns_false_when_not_in_failed_dict(self, plugin_manager):
|
||||
"""reload_failed_plugin returns (False, msg) when dir not in failed dict."""
|
||||
result = await plugin_manager.reload_failed_plugin("unknown")
|
||||
assert result == (False, "插件不存在于失败列表中")
|
||||
@@ -0,0 +1,262 @@
|
||||
"""Unit tests for astrbot.core.utils.storage_cleaner.
|
||||
|
||||
Uses tmp_path for filesystem-level assertions.
|
||||
"""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.utils.storage_cleaner import StorageCleaner
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_file(path: Path, size: int) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_bytes(b"x" * size)
|
||||
|
||||
|
||||
def _assert_exists(path: Path, expected_size: int | None = None) -> None:
|
||||
assert path.exists(), f"Expected {path} to exist"
|
||||
if expected_size is not None:
|
||||
assert path.stat().st_size == expected_size
|
||||
|
||||
|
||||
def _assert_missing(path: Path) -> None:
|
||||
assert not path.exists(), f"Expected {path} to be removed"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStorageCleanerStatus:
|
||||
def test_get_status_reports_logs_and_cache(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
temp = data / "temp"
|
||||
logs = data / "logs"
|
||||
|
||||
_make_file(temp / "x.wav", 100)
|
||||
_make_file(data / "plugins.json", 50)
|
||||
_make_file(logs / "astrbot.log", 200)
|
||||
_make_file(logs / "astrbot.2026-01-01.log", 80)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": True, "log_file_path": "logs/astrbot.log"},
|
||||
data_dir=data,
|
||||
temp_dir=temp,
|
||||
)
|
||||
|
||||
status = cleaner.get_status()
|
||||
|
||||
assert status["logs"]["size_bytes"] == 280 # 200 + 80
|
||||
assert status["logs"]["file_count"] == 2
|
||||
assert status["cache"]["size_bytes"] == 150 # 100 + 50
|
||||
assert status["cache"]["file_count"] == 2
|
||||
assert status["total_bytes"] == 430
|
||||
|
||||
def test_get_status_when_dirs_missing(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": False},
|
||||
data_dir=data,
|
||||
temp_dir=data / "temp",
|
||||
)
|
||||
status = cleaner.get_status()
|
||||
assert status["logs"]["size_bytes"] == 0
|
||||
assert status["logs"]["file_count"] == 0
|
||||
assert status["cache"]["size_bytes"] == 0
|
||||
assert status["total_bytes"] == 0
|
||||
|
||||
def test_get_status_includes_cache_extra_files(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
temp = data / "temp"
|
||||
logs = data / "logs"
|
||||
_make_file(logs / "astrbot.log", 10)
|
||||
_make_file(temp / "a.bin", 10)
|
||||
|
||||
_make_file(data / "plugins_custom_foo.json", 20)
|
||||
_make_file(data / "sandbox_skills_cache.json", 30)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": True},
|
||||
data_dir=data,
|
||||
temp_dir=temp,
|
||||
)
|
||||
status = cleaner.get_status()
|
||||
assert status["cache"]["file_count"] >= 3 # temp + custom* + sandbox
|
||||
assert status["cache"]["size_bytes"] >= 60
|
||||
|
||||
|
||||
class TestStorageCleanerCleanup:
|
||||
def test_cleanup_all_removes_logs_and_cache(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
temp = data / "temp"
|
||||
logs = data / "logs"
|
||||
|
||||
active_log = logs / "astrbot.log"
|
||||
rotated_log = logs / "astrbot.2026-03-01.log"
|
||||
_make_file(active_log, 100)
|
||||
_make_file(rotated_log, 50)
|
||||
_make_file(temp / "t.bin", 200)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": True, "log_file_path": "logs/astrbot.log"},
|
||||
data_dir=data,
|
||||
temp_dir=temp,
|
||||
)
|
||||
result = cleaner.cleanup("all")
|
||||
|
||||
# Active log → truncated
|
||||
_assert_exists(active_log, 0)
|
||||
# Rotated log → deleted
|
||||
_assert_missing(rotated_log)
|
||||
# Temp cache → deleted
|
||||
_assert_missing(temp / "t.bin")
|
||||
|
||||
assert result["removed_bytes"] == 350
|
||||
assert result["processed_files"] == 3
|
||||
assert result["deleted_files"] == 2
|
||||
assert result["truncated_files"] == 1
|
||||
|
||||
def test_cleanup_logs_target_only(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
logs = data / "logs"
|
||||
temp = data / "temp"
|
||||
|
||||
_make_file(logs / "astrbot.log", 80)
|
||||
_make_file(temp / "c.bin", 999)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": True, "log_file_path": "logs/astrbot.log"},
|
||||
data_dir=data,
|
||||
temp_dir=temp,
|
||||
)
|
||||
result = cleaner.cleanup("logs")
|
||||
|
||||
_assert_exists(logs / "astrbot.log", 0)
|
||||
# temp should be untouched
|
||||
_assert_exists(temp / "c.bin", 999)
|
||||
assert result["removed_bytes"] == 80
|
||||
|
||||
def test_cleanup_cache_target_only(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
logs = data / "logs"
|
||||
temp = data / "temp"
|
||||
|
||||
_make_file(logs / "astrbot.log", 80)
|
||||
_make_file(temp / "c.bin", 200)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": True, "log_file_path": "logs/astrbot.log"},
|
||||
data_dir=data,
|
||||
temp_dir=temp,
|
||||
)
|
||||
result = cleaner.cleanup("cache")
|
||||
|
||||
_assert_exists(logs / "astrbot.log", 80) # untouched
|
||||
_assert_missing(temp / "c.bin")
|
||||
assert result["removed_bytes"] == 200
|
||||
|
||||
def test_cleanup_invalid_target_raises(self, tmp_path: Path):
|
||||
cleaner = StorageCleaner({}, data_dir=tmp_path, temp_dir=tmp_path)
|
||||
with pytest.raises(ValueError, match="Unsupported cleanup target"):
|
||||
cleaner.cleanup("invalid")
|
||||
|
||||
def test_cleanup_with_no_config(self, tmp_path: Path):
|
||||
"""All config options off; no active log files, so everything gets deleted."""
|
||||
data = tmp_path / "data"
|
||||
_make_file(data / "logs" / "astrbot.log", 50)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": False},
|
||||
data_dir=data,
|
||||
temp_dir=data / "temp",
|
||||
)
|
||||
result = cleaner.cleanup("logs")
|
||||
# With log_file_enable=False, astrbot.log is NOT in active_log_files,
|
||||
# so it gets deleted, not truncated.
|
||||
_assert_missing(data / "logs" / "astrbot.log")
|
||||
assert result["deleted_files"] == 1
|
||||
assert result["truncated_files"] == 0
|
||||
|
||||
|
||||
class TestStorageCleanerEdgeCases:
|
||||
def test_cleanup_removes_empty_temp_dirs(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
temp = data / "temp"
|
||||
nested = temp / "sub" / "nested"
|
||||
_make_file(nested / "f.bin", 100)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{},
|
||||
data_dir=data,
|
||||
temp_dir=temp,
|
||||
)
|
||||
cleaner.cleanup("cache")
|
||||
|
||||
assert not nested.exists()
|
||||
assert temp.exists() # root temp should be recreated by _cleanup_target
|
||||
|
||||
def test_cleanup_skips_inexistent_file(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": False},
|
||||
data_dir=data,
|
||||
temp_dir=data / "temp",
|
||||
)
|
||||
# No files at all – should not crash.
|
||||
result = cleaner.cleanup("all")
|
||||
assert result["failed_files"] == 0
|
||||
|
||||
def test_cleanup_handles_stat_os_error(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
logs = data / "logs"
|
||||
_make_file(logs / "astrbot.log", 100)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": True, "log_file_path": "logs/astrbot.log"},
|
||||
data_dir=data,
|
||||
temp_dir=data / "temp",
|
||||
)
|
||||
|
||||
original_stat = (logs / "astrbot.log").stat
|
||||
|
||||
def _broken_stat():
|
||||
raise OSError(13, "Permission denied")
|
||||
|
||||
(logs / "astrbot.log").stat = _broken_stat # type: ignore[method-assign]
|
||||
|
||||
result = cleaner.cleanup("logs")
|
||||
assert result["failed_files"] == 1
|
||||
|
||||
(logs / "astrbot.log").stat = original_stat
|
||||
|
||||
def test_cleanup_handles_unlink_os_error(self, tmp_path: Path):
|
||||
data = tmp_path / "data"
|
||||
logs = data / "logs"
|
||||
_make_file(logs / "astrbot.log", 100)
|
||||
|
||||
cleaner = StorageCleaner(
|
||||
{"log_file_enable": False},
|
||||
data_dir=data,
|
||||
temp_dir=data / "temp",
|
||||
)
|
||||
|
||||
original_unlink = (logs / "astrbot.log").unlink
|
||||
|
||||
def _broken_unlink():
|
||||
raise OSError(13, "Permission denied")
|
||||
|
||||
(logs / "astrbot.log").unlink = _broken_unlink # type: ignore[method-assign]
|
||||
|
||||
result = cleaner.cleanup("logs")
|
||||
assert result["failed_files"] == 1
|
||||
|
||||
(logs / "astrbot.log").unlink = original_unlink
|
||||
@@ -0,0 +1,259 @@
|
||||
"""Unit tests for astrbot.core.utils.temp_dir_cleaner.
|
||||
|
||||
Covers parse_size_to_bytes, cleanup_once, async lifecycle, and error paths.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.utils.temp_dir_cleaner import TempDirCleaner, parse_size_to_bytes
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# parse_size_to_bytes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestParseSizeToBytes:
|
||||
def test_valid_mb_string(self):
|
||||
assert parse_size_to_bytes("1024") == 1024 * 1024**2
|
||||
|
||||
def test_valid_float_mb(self):
|
||||
assert parse_size_to_bytes(0.5) == int(0.5 * 1024**2)
|
||||
|
||||
def test_zero_returns_zero(self):
|
||||
assert parse_size_to_bytes(0) == 0
|
||||
|
||||
def test_none_returns_zero(self):
|
||||
assert parse_size_to_bytes(None) == 0
|
||||
|
||||
def test_invalid_string_returns_zero(self):
|
||||
assert parse_size_to_bytes("not-a-number") == 0
|
||||
|
||||
def test_negative_value_returns_zero(self):
|
||||
assert parse_size_to_bytes("-10") == 0
|
||||
|
||||
def test_whitespace_string(self):
|
||||
assert parse_size_to_bytes(" 512 ") == 512 * 1024**2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# helpers for filesystem tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _write_file(path: Path, size: int, mtime: float | None = None) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_bytes(b"x" * size)
|
||||
if mtime is not None:
|
||||
os.utime(path, (mtime, mtime))
|
||||
|
||||
|
||||
def _file_sizes(temp_dir: Path) -> list[tuple[Path, int]]:
|
||||
return [
|
||||
(f, f.stat().st_size)
|
||||
for f in sorted(temp_dir.rglob("*"))
|
||||
if f.is_file()
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# cleanup_once
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCleanupOnce:
|
||||
def test_noop_when_below_limit(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "temp"
|
||||
temp_dir.mkdir(parents=True, exist_ok=True)
|
||||
_write_file(temp_dir / "a.bin", 100, time.time())
|
||||
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "1", # 1 MB → effectively unlimited
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
cleaner.cleanup_once()
|
||||
|
||||
files = _file_sizes(temp_dir)
|
||||
assert len(files) == 1
|
||||
assert files[0][1] == 100
|
||||
|
||||
def test_removes_oldest_files_when_over_limit(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "temp"
|
||||
temp_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
base = time.time() - 1000
|
||||
_write_file(temp_dir / "old.bin", 400, base)
|
||||
_write_file(temp_dir / "mid.bin", 300, base + 10)
|
||||
_write_file(temp_dir / "new.bin", 300, base + 20)
|
||||
|
||||
# Limit = 0.0008 MB ≈ 838 bytes; total = 1000 => release 30% = 300
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "0.0008",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
cleaner.cleanup_once()
|
||||
|
||||
remaining = _file_sizes(temp_dir)
|
||||
remaining_total = sum(sz for _, sz in remaining)
|
||||
# old.bin (oldest) should have been deleted first
|
||||
assert (temp_dir / "old.bin").exists() is False
|
||||
# remaining <= 700 (total - minimum 30%)
|
||||
assert remaining_total <= 700
|
||||
|
||||
def test_noop_when_temp_dir_missing(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "nonexistent"
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "0.001",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
# Should not raise
|
||||
cleaner.cleanup_once()
|
||||
|
||||
def test_handles_unlink_os_error_gracefully(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "temp"
|
||||
temp_dir.mkdir(parents=True, exist_ok=True)
|
||||
_write_file(temp_dir / "bad.bin", 9999, time.time() - 100)
|
||||
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "0.001",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
|
||||
def _broken_unlink():
|
||||
raise OSError(13, "Permission denied")
|
||||
|
||||
(temp_dir / "bad.bin").unlink = _broken_unlink # type: ignore[method-assign]
|
||||
|
||||
# Should log warning but not crash
|
||||
with patch("astrbot.core.utils.temp_dir_cleaner.logger.warning") as mock_warn:
|
||||
cleaner.cleanup_once()
|
||||
mock_warn.assert_called()
|
||||
assert any(
|
||||
"Permission denied" in str(c) for c in mock_warn.call_args_list
|
||||
)
|
||||
|
||||
def test_invalid_config_falls_back_to_default(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "temp"
|
||||
temp_dir.mkdir(parents=True, exist_ok=True)
|
||||
_write_file(temp_dir / "a.bin", 50, time.time())
|
||||
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "invalid",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
with patch("astrbot.core.utils.temp_dir_cleaner.logger.warning") as mock_warn:
|
||||
cleaner.cleanup_once()
|
||||
mock_warn.assert_called_once()
|
||||
assert "fallback" in str(mock_warn.call_args[0][0])
|
||||
|
||||
def test_cleanup_once_removes_empty_dirs(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "temp"
|
||||
nested = temp_dir / "a" / "b"
|
||||
_write_file(nested / "f.bin", 1000000, time.time() - 5000)
|
||||
# Also create an already-empty subdir
|
||||
empty_dir = temp_dir / "empty"
|
||||
empty_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "0.001",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
cleaner.cleanup_once()
|
||||
|
||||
# The file in nested is old → gets deleted, dirs cleaned up
|
||||
assert not nested.exists()
|
||||
# empty dir should be removed by _cleanup_empty_dirs
|
||||
_assert_dir_gone = not empty_dir.exists()
|
||||
# either nested or empty (or both) should be gone
|
||||
assert _assert_dir_gone or not nested.exists()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# scan / async lifecycle
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestScanAndAsyncLifecycle:
|
||||
def test_scan_empty_dir_returns_zeroes(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "empty_temp"
|
||||
temp_dir.mkdir(parents=True, exist_ok=True)
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "100",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
total, files = cleaner._scan_temp_files()
|
||||
assert total == 0
|
||||
assert files == []
|
||||
|
||||
def test_scan_missing_dir_returns_zeroes(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "ghost"
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "100",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
total, files = cleaner._scan_temp_files()
|
||||
assert total == 0
|
||||
assert files == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_sets_event(self):
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "100",
|
||||
)
|
||||
await cleaner.stop()
|
||||
assert cleaner._stop_event.is_set() is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_stops_when_stop_is_called(self, tmp_path: Path):
|
||||
temp_dir = tmp_path / "temp"
|
||||
temp_dir.mkdir(parents=True, exist_ok=True)
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "1000",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
|
||||
task = asyncio.create_task(cleaner.run())
|
||||
await asyncio.sleep(0.05)
|
||||
assert task.done() is False
|
||||
|
||||
await cleaner.stop()
|
||||
await asyncio.wait_for(task, timeout=2.0)
|
||||
assert task.done() is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_logs_exception_and_continues(self, tmp_path: Path):
|
||||
"""When cleanup_once raises, run() should log and continue the loop."""
|
||||
temp_dir = tmp_path / "temp"
|
||||
temp_dir.mkdir(parents=True, exist_ok=True)
|
||||
cleaner = TempDirCleaner(
|
||||
max_size_getter=lambda: "1000",
|
||||
temp_dir=temp_dir,
|
||||
)
|
||||
|
||||
# Make cleanup_once raise once, then succeed
|
||||
call_count = 0
|
||||
|
||||
def _flaky_cleanup():
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
raise RuntimeError("transient failure")
|
||||
|
||||
cleaner.cleanup_once = _flaky_cleanup # type: ignore[method-assign]
|
||||
|
||||
with patch(
|
||||
"astrbot.core.utils.temp_dir_cleaner.logger.error"
|
||||
) as mock_log_error:
|
||||
task = asyncio.create_task(cleaner.run())
|
||||
await asyncio.sleep(0.05)
|
||||
await cleaner.stop()
|
||||
await asyncio.wait_for(task, timeout=2.0)
|
||||
|
||||
mock_log_error.assert_called()
|
||||
error_msg = str(mock_log_error.call_args[0][0])
|
||||
assert "transient failure" in error_msg or "failed" in error_msg
|
||||
@@ -0,0 +1,344 @@
|
||||
"""Unit tests for astrbot.core.agent.tool: FunctionTool, ToolSchema, ToolSet."""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import MagicMock, AsyncMock
|
||||
from pydantic import ValidationError
|
||||
|
||||
from astrbot.core.agent.tool import (
|
||||
FunctionTool,
|
||||
ToolSchema,
|
||||
ToolSet,
|
||||
ToolArgumentSpec,
|
||||
ToolExecResult,
|
||||
)
|
||||
|
||||
|
||||
class TestToolSchema:
|
||||
"""ToolSchema construction and validation."""
|
||||
|
||||
def test_tool_schema_minimal(self):
|
||||
"""Construct with only name and description."""
|
||||
schema = ToolSchema(name="test_tool", description="A test tool")
|
||||
assert schema.name == "test_tool"
|
||||
assert schema.description == "A test tool"
|
||||
assert schema.parameters is None
|
||||
assert schema.active is True
|
||||
|
||||
def test_tool_schema_with_parameters(self):
|
||||
"""Construct with valid JSON Schema parameters."""
|
||||
params = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {"type": "string", "description": "Search query"}
|
||||
},
|
||||
"required": ["query"],
|
||||
}
|
||||
schema = ToolSchema(name="search", description="Search tool", parameters=params)
|
||||
assert schema.name == "search"
|
||||
assert schema.parameters == params
|
||||
|
||||
def test_tool_schema_invalid_parameters_raises(self):
|
||||
"""Reject parameters that are not valid JSON Schema."""
|
||||
with pytest.raises(ValidationError):
|
||||
ToolSchema(
|
||||
name="bad",
|
||||
description="Bad schema",
|
||||
parameters={"type": "nonexistent"},
|
||||
)
|
||||
|
||||
def test_tool_schema_inactive(self):
|
||||
"""Explicitly set active=False."""
|
||||
schema = ToolSchema(name="off", description="Disabled", active=False)
|
||||
assert schema.active is False
|
||||
|
||||
def test_tool_schema_parameters_none_allowed(self):
|
||||
"""parameters=None is valid and passes validation."""
|
||||
schema = ToolSchema(name="no_params", description="No parameters", parameters=None)
|
||||
assert schema.parameters is None
|
||||
|
||||
def test_tool_schema_empty_parameters_object(self):
|
||||
"""Empty object parameters are valid JSON Schema."""
|
||||
schema = ToolSchema(
|
||||
name="empty",
|
||||
description="Empty params",
|
||||
parameters={"type": "object", "properties": {}},
|
||||
)
|
||||
assert schema.parameters == {"type": "object", "properties": {}}
|
||||
|
||||
|
||||
class TestFunctionTool:
|
||||
"""FunctionTool construction and call interface."""
|
||||
|
||||
def test_function_tool_minimal(self):
|
||||
"""Construct with minimal fields."""
|
||||
tool = FunctionTool(name="echo", description="Echo tool")
|
||||
assert tool.name == "echo"
|
||||
assert tool.description == "Echo tool"
|
||||
assert tool.handler is None
|
||||
assert tool.active is True
|
||||
assert tool.is_background_task is False
|
||||
assert tool.source == "plugin"
|
||||
|
||||
def test_function_tool_with_handler(self):
|
||||
"""Construct with an async handler."""
|
||||
async def handler(context, **kwargs):
|
||||
return "ok"
|
||||
tool = FunctionTool(
|
||||
name="greet",
|
||||
description="Greet",
|
||||
handler=handler,
|
||||
handler_module_path="tests.test_tool",
|
||||
)
|
||||
assert tool.handler is handler
|
||||
assert tool.handler_module_path == "tests.test_tool"
|
||||
|
||||
def test_function_tool_background_task(self):
|
||||
"""Construct with is_background_task=True."""
|
||||
tool = FunctionTool(
|
||||
name="bg",
|
||||
description="Background",
|
||||
is_background_task=True,
|
||||
)
|
||||
assert tool.is_background_task is True
|
||||
|
||||
def test_function_tool_source_values(self):
|
||||
"""Construct with different source values."""
|
||||
for source in ("plugin", "internal", "mcp"):
|
||||
tool = FunctionTool(
|
||||
name=f"tool_{source}",
|
||||
description=source,
|
||||
source=source,
|
||||
)
|
||||
assert tool.source == source
|
||||
|
||||
def test_function_tool_repr(self):
|
||||
"""__repr__ returns a meaningful string."""
|
||||
tool = FunctionTool(name="sum", description="Sum numbers")
|
||||
rep = repr(tool)
|
||||
assert "FuncTool" in rep
|
||||
assert "sum" in rep
|
||||
|
||||
def test_function_tool_call_not_implemented(self):
|
||||
"""call() raises NotImplementedError when no handler is set."""
|
||||
tool = FunctionTool(name="todo", description="Not implemented")
|
||||
with pytest.raises(NotImplementedError, match="FunctionTool.call"):
|
||||
import asyncio
|
||||
asyncio.run(tool.call(MagicMock()))
|
||||
|
||||
def test_function_tool_with_parameters(self):
|
||||
"""Construct with valid parameters."""
|
||||
params = {
|
||||
"type": "object",
|
||||
"properties": {"x": {"type": "integer"}},
|
||||
}
|
||||
tool = FunctionTool(name="add", description="Add", parameters=params)
|
||||
assert tool.parameters == params
|
||||
|
||||
def test_function_tool_active_false(self):
|
||||
"""Construct with active=False."""
|
||||
tool = FunctionTool(name="inactive_tool", description="Not active", active=False)
|
||||
assert tool.active is False
|
||||
|
||||
def test_function_tool_inherits_validation(self):
|
||||
"""FunctionTool inherits ToolSchema parameter validation."""
|
||||
with pytest.raises(ValidationError):
|
||||
FunctionTool(
|
||||
name="bad",
|
||||
description="Bad",
|
||||
parameters={"type": "madeup"},
|
||||
)
|
||||
|
||||
|
||||
class TestToolSet:
|
||||
"""ToolSet add/remove/get and serialization methods."""
|
||||
|
||||
def test_empty_toolset(self):
|
||||
"""New ToolSet is empty."""
|
||||
ts = ToolSet()
|
||||
assert ts.empty() is True
|
||||
assert len(ts) == 0
|
||||
assert bool(ts) is False
|
||||
|
||||
def test_add_tool(self):
|
||||
"""Add a tool increases length."""
|
||||
ts = ToolSet()
|
||||
tool = FunctionTool(name="a", description="Tool A")
|
||||
ts.add_tool(tool)
|
||||
assert len(ts) == 1
|
||||
assert ts.empty() is False
|
||||
assert bool(ts) is True
|
||||
|
||||
def test_add_duplicate_name_overwrites(self):
|
||||
"""Adding a tool with same name replaces the existing one."""
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="dup", description="Original"))
|
||||
ts.add_tool(FunctionTool(name="dup", description="Replacement"))
|
||||
assert len(ts) == 1
|
||||
tool = ts.get_tool("dup")
|
||||
assert tool is not None
|
||||
assert tool.description == "Replacement"
|
||||
|
||||
def test_add_duplicate_active_prefers_active(self):
|
||||
"""When a duplicate is added, active=True wins over active=False."""
|
||||
ts = ToolSet()
|
||||
inactive = FunctionTool(name="x", description="Inactive", active=False)
|
||||
active = FunctionTool(name="x", description="Active", active=True)
|
||||
ts.add_tool(inactive)
|
||||
ts.add_tool(active)
|
||||
tool = ts.get_tool("x")
|
||||
assert tool is not None
|
||||
assert tool.description == "Active"
|
||||
|
||||
def test_add_duplicate_inactive_does_not_replace_active(self):
|
||||
"""Adding an inactive tool with an already active name does not overwrite."""
|
||||
ts = ToolSet()
|
||||
active = FunctionTool(name="y", description="Active", active=True)
|
||||
inactive = FunctionTool(name="y", description="Inactive", active=False)
|
||||
ts.add_tool(active)
|
||||
ts.add_tool(inactive)
|
||||
tool = ts.get_tool("y")
|
||||
assert tool is not None
|
||||
assert tool.description == "Active"
|
||||
|
||||
def test_remove_tool(self):
|
||||
"""Remove a tool by name."""
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="keep", description="Keep"))
|
||||
ts.add_tool(FunctionTool(name="remove_me", description="Remove"))
|
||||
ts.remove_tool("remove_me")
|
||||
assert len(ts) == 1
|
||||
assert ts.get_tool("keep") is not None
|
||||
assert ts.get_tool("remove_me") is None
|
||||
|
||||
def test_remove_nonexistent_tool(self):
|
||||
"""Removing a non-existent tool does not raise."""
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="a", description="A"))
|
||||
ts.remove_tool("nonexistent")
|
||||
assert len(ts) == 1
|
||||
|
||||
def test_get_tool_returns_none_for_non_functiontool(self):
|
||||
"""get_tool returns None when a ToolSchema (not FunctionTool) exists with that name."""
|
||||
ts = ToolSet()
|
||||
schema = ToolSchema(name="plain", description="Plain schema")
|
||||
ts.add_tool(schema)
|
||||
assert ts.get_tool("plain") is None
|
||||
|
||||
def test_get_tool_returns_functiontool(self):
|
||||
"""get_tool returns the FunctionTool when present."""
|
||||
ts = ToolSet()
|
||||
tool = FunctionTool(name="found", description="Found tool")
|
||||
ts.add_tool(tool)
|
||||
result = ts.get_tool("found")
|
||||
assert result is tool
|
||||
|
||||
def test_normalize_sorts_by_name(self):
|
||||
"""normalize() sorts tools alphabetically by name."""
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="z", description="Z"))
|
||||
ts.add_tool(FunctionTool(name="a", description="A"))
|
||||
ts.add_tool(FunctionTool(name="m", description="M"))
|
||||
ts.normalize()
|
||||
assert [t.name for t in ts.tools] == ["a", "m", "z"]
|
||||
|
||||
def test_names(self):
|
||||
"""names() returns list of tool names."""
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="a", description="A"))
|
||||
ts.add_tool(FunctionTool(name="b", description="B"))
|
||||
assert ts.names() == ["a", "b"]
|
||||
|
||||
def test_func_list(self):
|
||||
"""func_list only includes FunctionTool instances."""
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="func", description="Func"))
|
||||
ts.add_tool(ToolSchema(name="plain", description="Plain"))
|
||||
assert len(ts.func_list) == 1
|
||||
assert ts.func_list[0].name == "func"
|
||||
|
||||
def test_merge(self):
|
||||
"""merge() combines tools from another ToolSet."""
|
||||
ts1 = ToolSet()
|
||||
ts1.add_tool(FunctionTool(name="a", description="A"))
|
||||
ts2 = ToolSet()
|
||||
ts2.add_tool(FunctionTool(name="b", description="B"))
|
||||
ts1.merge(ts2)
|
||||
assert len(ts1) == 2
|
||||
|
||||
def test_iteration(self):
|
||||
"""ToolSet is iterable and yields ToolSchema items."""
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="a", description="A"))
|
||||
ts.add_tool(FunctionTool(name="b", description="B"))
|
||||
names = [t.name for t in ts]
|
||||
assert names == ["a", "b"]
|
||||
|
||||
def test_get_light_tool_set(self):
|
||||
"""get_light_tool_set returns tools with empty parameters and no handler."""
|
||||
params = {
|
||||
"type": "object",
|
||||
"properties": {"q": {"type": "string"}},
|
||||
}
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(
|
||||
name="search", description="Search", parameters=params,
|
||||
handler=AsyncMock(),
|
||||
))
|
||||
light = ts.get_light_tool_set()
|
||||
assert len(light) == 1
|
||||
light_tool = light.get_tool("search")
|
||||
assert light_tool is not None
|
||||
assert light_tool.description == "Search"
|
||||
assert light_tool.parameters == {"type": "object", "properties": {}}
|
||||
assert light_tool.handler is None
|
||||
|
||||
def test_get_light_tool_skips_inactive(self):
|
||||
"""get_light_tool_set skips inactive tools."""
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="active", description="Active", active=True))
|
||||
ts.add_tool(FunctionTool(name="inactive", description="Inactive", active=False))
|
||||
light = ts.get_light_tool_set()
|
||||
assert light.names() == ["active"]
|
||||
|
||||
def test_get_param_only_tool_set(self):
|
||||
"""get_param_only_tool_set returns tools with empty description."""
|
||||
params = {
|
||||
"type": "object",
|
||||
"properties": {"x": {"type": "integer"}},
|
||||
}
|
||||
ts = ToolSet()
|
||||
ts.add_tool(FunctionTool(name="calc", description="Calc", parameters=params))
|
||||
param_only = ts.get_param_only_tool_set()
|
||||
assert len(param_only) == 1
|
||||
tool = param_only.get_tool("calc")
|
||||
assert tool is not None
|
||||
assert tool.description == ""
|
||||
assert tool.parameters == params
|
||||
|
||||
def test_anthropic_schema(self):
|
||||
"""anthropic_schema returns Anthropic-compatible format."""
|
||||
ts = ToolSet()
|
||||
params = {
|
||||
"type": "object",
|
||||
"properties": {"x": {"type": "integer"}},
|
||||
"required": ["x"],
|
||||
}
|
||||
ts.add_tool(FunctionTool(name="add", description="Add", parameters=params))
|
||||
schema = ts.anthropic_schema()
|
||||
assert len(schema) == 1
|
||||
assert schema[0]["name"] == "add"
|
||||
assert schema[0]["description"] == "Add"
|
||||
assert schema[0]["input_schema"]["properties"] == {"x": {"type": "integer"}}
|
||||
assert schema[0]["input_schema"]["required"] == ["x"]
|
||||
|
||||
def test_google_schema(self):
|
||||
"""google_schema returns Google GenAI-compatible format."""
|
||||
ts = ToolSet()
|
||||
params = {
|
||||
"type": "object",
|
||||
"properties": {"val": {"type": "number", "description": "A value"}},
|
||||
}
|
||||
ts.add_tool(FunctionTool(name="sqrt", description="Square root", parameters=params))
|
||||
schema = ts.google_schema()
|
||||
assert "function_declarations" in schema
|
||||
assert schema["function_declarations"][0]["name"] == "sqrt"
|
||||
@@ -0,0 +1,136 @@
|
||||
"""Unit tests for astrbot.core.agent.tool_executor: BaseFunctionToolExecutor abstract signature."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.agent.run_context import ContextWrapper
|
||||
from astrbot.core.agent.tool import FunctionTool
|
||||
from astrbot.core.agent.tool_executor import BaseFunctionToolExecutor
|
||||
|
||||
|
||||
class TestBaseFunctionToolExecutor:
|
||||
"""BaseFunctionToolExecutor cannot be instantiated; metachecks on the ABC."""
|
||||
|
||||
def test_cannot_instantiate_abstract(self):
|
||||
"""BaseFunctionToolExecutor cannot be instantiated directly (abstract)."""
|
||||
with pytest.raises(TypeError, match="abstract"):
|
||||
BaseFunctionToolExecutor()
|
||||
|
||||
def test_execute_is_abstract_classmethod(self):
|
||||
"""execute is both a classmethod and abstractmethod."""
|
||||
assert hasattr(BaseFunctionToolExecutor.execute, "__isabstractmethod__")
|
||||
assert BaseFunctionToolExecutor.execute.__isabstractmethod__ is True
|
||||
|
||||
def test_execute_is_classmethod_descriptor(self):
|
||||
"""execute is a classmethod (has __func__)."""
|
||||
assert hasattr(BaseFunctionToolExecutor.execute, "__func__")
|
||||
|
||||
def test_concrete_subclass_must_implement_execute(self):
|
||||
"""A subclass without execute is still abstract."""
|
||||
class Missing(BaseFunctionToolExecutor):
|
||||
pass
|
||||
with pytest.raises(TypeError, match="abstract"):
|
||||
Missing()
|
||||
|
||||
def test_concrete_subclass_with_execute_can_instantiate(self):
|
||||
"""A subclass that implements execute can be instantiated."""
|
||||
class Concrete(BaseFunctionToolExecutor):
|
||||
@classmethod
|
||||
async def execute(cls, tool, run_context, **tool_args):
|
||||
yield "result"
|
||||
|
||||
instance = Concrete()
|
||||
assert isinstance(instance, BaseFunctionToolExecutor)
|
||||
|
||||
def test_execute_signature_matches(self):
|
||||
"""execute has the expected parameter names."""
|
||||
import inspect
|
||||
sig = inspect.signature(
|
||||
BaseFunctionToolExecutor.__dict__["execute"].__func__
|
||||
)
|
||||
param_names = list(sig.parameters.keys())
|
||||
assert "tool" in param_names
|
||||
assert "run_context" in param_names
|
||||
|
||||
def test_execute_tool_parameter_type_hint(self):
|
||||
"""The tool parameter is annotated as FunctionTool."""
|
||||
import inspect
|
||||
sig = inspect.signature(
|
||||
BaseFunctionToolExecutor.__dict__["execute"].__func__
|
||||
)
|
||||
tool_param = sig.parameters["tool"]
|
||||
assert tool_param.annotation is FunctionTool
|
||||
|
||||
def test_execute_run_context_type_hint(self):
|
||||
"""The run_context parameter is annotated as ContextWrapper."""
|
||||
import inspect
|
||||
sig = inspect.signature(
|
||||
BaseFunctionToolExecutor.__dict__["execute"].__func__
|
||||
)
|
||||
ctx_param = sig.parameters["run_context"]
|
||||
origin = getattr(ctx_param.annotation, "__origin__", None)
|
||||
assert origin is ContextWrapper
|
||||
|
||||
def test_execute_returns_async_generator(self):
|
||||
"""execute return annotation is AsyncGenerator."""
|
||||
import inspect
|
||||
sig = inspect.signature(
|
||||
BaseFunctionToolExecutor.__dict__["execute"].__func__
|
||||
)
|
||||
return_annotation = sig.return_annotation
|
||||
assert "AsyncGenerator" in str(return_annotation)
|
||||
|
||||
def test_concrete_subclass_execute_yields(self):
|
||||
"""A concrete subclass can actually yield values."""
|
||||
class Tester(BaseFunctionToolExecutor):
|
||||
@classmethod
|
||||
async def execute(cls, tool, run_context, **tool_args):
|
||||
yield "step1"
|
||||
yield "step2"
|
||||
|
||||
import asyncio
|
||||
tool = FunctionTool(name="test", description="test")
|
||||
ctx = ContextWrapper(context="test")
|
||||
gen = Tester.execute(tool, ctx)
|
||||
results = asyncio.run(async_collect(gen))
|
||||
assert results == ["step1", "step2"]
|
||||
|
||||
def test_concrete_subclass_generic_parameter(self):
|
||||
"""Subclass can specialize the generic type parameter."""
|
||||
class TypedExecutor(BaseFunctionToolExecutor[str]):
|
||||
@classmethod
|
||||
async def execute(cls, tool, run_context, **tool_args):
|
||||
yield run_context.context
|
||||
|
||||
import asyncio
|
||||
tool = FunctionTool(name="t", description="t")
|
||||
ctx = ContextWrapper(context="hello_generic")
|
||||
gen = TypedExecutor.execute(tool, ctx)
|
||||
results = asyncio.run(async_collect(gen))
|
||||
assert results == ["hello_generic"]
|
||||
|
||||
def test_subclass_with_kwargs_passthrough(self):
|
||||
"""execute passes **tool_args through."""
|
||||
class KwargsExecutor(BaseFunctionToolExecutor):
|
||||
@classmethod
|
||||
async def execute(cls, tool, run_context, **tool_args):
|
||||
yield tool_args
|
||||
|
||||
import asyncio
|
||||
tool = FunctionTool(name="t", description="t")
|
||||
ctx = ContextWrapper(context="ctx")
|
||||
gen = KwargsExecutor.execute(tool, ctx, x=1, y="two")
|
||||
results = asyncio.run(async_collect(gen))
|
||||
assert results[0] == {"x": 1, "y": "two"}
|
||||
|
||||
|
||||
async def async_collect(async_gen):
|
||||
"""Collect all items from an async generator into a list."""
|
||||
results = []
|
||||
async for item in async_gen:
|
||||
results.append(item)
|
||||
return results
|
||||
Reference in New Issue
Block a user