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:
LIghtJUNction
2026-04-29 08:40:26 +08:00
parent daee102c07
commit 169a856a3a
39 changed files with 14299 additions and 9 deletions
+280
View File
@@ -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
+323
View File
@@ -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
+265
View File
@@ -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"
+267
View File
@@ -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)
+97
View File
@@ -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
+529
View File
@@ -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"]
+785
View File
@@ -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")
+242
View File
@@ -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()
+428
View File
@@ -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
+254
View File
@@ -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
+454
View File
@@ -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 == []
+211
View File
@@ -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
+396
View File
@@ -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
+267
View File
@@ -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()
+314
View File
@@ -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)
+926
View File
@@ -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"]
+657
View File
@@ -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,
)
+283
View File
@@ -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
+335
View File
@@ -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)
+169
View File
@@ -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
+191 -9
View File
@@ -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()
+238
View File
@@ -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)
+346
View File
@@ -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
+530
View File
@@ -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()
+400
View File
@@ -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)
+312
View File
@@ -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
+518
View File
@@ -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
+408
View File
@@ -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
+473
View File
@@ -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
+244
View File
@@ -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
+375
View File
@@ -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)
+331
View File
@@ -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()
+91
View File
@@ -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
+829
View File
@@ -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
+530
View File
@@ -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, "插件不存在于失败列表中")
+262
View File
@@ -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
+259
View File
@@ -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
+344
View File
@@ -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"
+136
View File
@@ -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