mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-06-10 17:35:57 +00:00
fix(skills): harden slash skill activation across chat channels (#3466)
* support slash skill activation * format slash skill activation * Preserve slash skill activation with uploads * Address slash skill review feedback * Address slash skill follow-up review * Fix lazy slash skill storage resolution * Keep slash skill activation out of system prompt * Address slash skill review issues * fix: harden slash skill command handling * feat(frontend): add slash skill autocomplete * fix: address slash skill review feedback * fix: preserve slash skill text for IM uploads
This commit is contained in:
@@ -21,6 +21,42 @@ from app.channels.message_bus import (
|
||||
ResolvedAttachment,
|
||||
)
|
||||
from app.channels.store import ChannelStore
|
||||
from deerflow.skills.types import Skill, SkillCategory
|
||||
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY
|
||||
|
||||
|
||||
def test_known_channel_command_detection_only_matches_control_commands():
|
||||
from app.channels.commands import is_known_channel_command
|
||||
|
||||
assert is_known_channel_command("/new")
|
||||
assert is_known_channel_command("/HELP now")
|
||||
assert not is_known_channel_command("/mnt/user-data/uploads/report.pdf")
|
||||
assert not is_known_channel_command("/data-analysis analyze uploads/foo.csv")
|
||||
assert not is_known_channel_command(" /new")
|
||||
|
||||
|
||||
def _make_channel_skill(tmp_path: Path, name: str, *, enabled: bool = True) -> Skill:
|
||||
skill_dir = tmp_path / name
|
||||
skill_dir.mkdir(parents=True, exist_ok=True)
|
||||
skill_file = skill_dir / "SKILL.md"
|
||||
skill_file.write_text(f"# {name}\n", encoding="utf-8")
|
||||
return Skill(
|
||||
name=name,
|
||||
description=f"Description for {name}",
|
||||
license="MIT",
|
||||
skill_dir=skill_dir,
|
||||
skill_file=skill_file,
|
||||
relative_path=Path(name),
|
||||
category=SkillCategory.CUSTOM,
|
||||
enabled=enabled,
|
||||
)
|
||||
|
||||
|
||||
def _make_channel_skill_storage(skills: list[Skill]):
|
||||
return SimpleNamespace(
|
||||
load_skills=lambda *, enabled_only: [skill for skill in skills if skill.enabled] if enabled_only else skills,
|
||||
get_container_root=lambda: "/mnt/skills",
|
||||
)
|
||||
|
||||
|
||||
def _run(coro):
|
||||
@@ -1334,6 +1370,496 @@ class TestChannelManager:
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_blank_text_is_reported_without_running_agent(self):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text=" ",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_not_called()
|
||||
assert outbound_received[0].text.startswith("Unknown command.")
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_rejects_multi_slash_control_command(self):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text="//help",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_not_called()
|
||||
assert outbound_received[0].text.startswith("Unknown command: //help.")
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_requires_control_command_at_start(self):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
|
||||
mock_client = _make_mock_langgraph_client(thread_id="new-thread-456")
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text=" /new",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.threads.create.assert_not_called()
|
||||
assert store.get_thread_id("test", "chat1") is None
|
||||
assert outbound_received[0].text.startswith("Unknown command: /new.")
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_outbound_thread_id_uses_topic_thread(self):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
store.set_thread_id("test", "chat1", "base-thread")
|
||||
store.set_thread_id("test", "chat1", "topic-thread", topic_id="topic-1")
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text="/status",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
topic_id="topic-1",
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
assert outbound_received[0].text == "Active thread: topic-thread"
|
||||
assert outbound_received[0].thread_id == "topic-thread"
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_slash_skill_routes_to_chat(self, tmp_path):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
manager._skill_storage = _make_channel_skill_storage([_make_channel_skill(tmp_path, "data-analysis")])
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text="/data-analysis analyze uploads/foo.csv",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_called_once()
|
||||
call_args = mock_client.runs.wait.call_args
|
||||
assert call_args[1]["input"]["messages"][0]["content"] == "/data-analysis analyze uploads/foo.csv"
|
||||
assert outbound_received[0].text == "Hello from agent!"
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_slash_skill_with_attachment_preserves_original_content(self, monkeypatch, tmp_path):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def fake_ingest(thread_id, msg):
|
||||
return [
|
||||
{
|
||||
"filename": "report.pdf",
|
||||
"size": 12,
|
||||
"path": "/mnt/user-data/uploads/report.pdf",
|
||||
"is_image": False,
|
||||
}
|
||||
]
|
||||
|
||||
monkeypatch.setattr("app.channels.manager._ingest_inbound_files", fake_ingest)
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
manager._skill_storage = _make_channel_skill_storage([_make_channel_skill(tmp_path, "data-analysis")])
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
original_text = "/data-analysis analyze report.pdf"
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text=original_text,
|
||||
files=[{"filename": "report.pdf"}],
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_called_once()
|
||||
human_message = mock_client.runs.wait.call_args[1]["input"]["messages"][0]
|
||||
assert human_message["content"].startswith("<uploaded_files>")
|
||||
assert original_text in human_message["content"]
|
||||
assert human_message["additional_kwargs"][ORIGINAL_USER_CONTENT_KEY] == original_text
|
||||
assert outbound_received[0].text == "Hello from agent!"
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_streaming_slash_skill_with_attachment_preserves_original_content(self, monkeypatch, tmp_path):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def fake_ingest(thread_id, msg):
|
||||
return [
|
||||
{
|
||||
"filename": "report.pdf",
|
||||
"size": 12,
|
||||
"path": "/mnt/user-data/uploads/report.pdf",
|
||||
"is_image": False,
|
||||
}
|
||||
]
|
||||
|
||||
monkeypatch.setattr("app.channels.manager._ingest_inbound_files", fake_ingest)
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
manager._skill_storage = _make_channel_skill_storage([_make_channel_skill(tmp_path, "data-analysis")])
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
mock_client.runs.stream = MagicMock(
|
||||
return_value=_make_async_iterator(
|
||||
[
|
||||
_make_stream_part(
|
||||
"values",
|
||||
{"messages": [{"type": "ai", "content": "streamed response"}]},
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
original_text = "/data-analysis analyze report.pdf"
|
||||
inbound = InboundMessage(
|
||||
channel_name="feishu",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text=original_text,
|
||||
files=[{"filename": "report.pdf"}],
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: any(message.is_final for message in outbound_received))
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.stream.assert_called_once()
|
||||
human_message = mock_client.runs.stream.call_args[1]["input"]["messages"][0]
|
||||
assert human_message["content"].startswith("<uploaded_files>")
|
||||
assert original_text in human_message["content"]
|
||||
assert human_message["additional_kwargs"][ORIGINAL_USER_CONTENT_KEY] == original_text
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_slash_skill_requires_command_at_start(self, tmp_path):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
manager._skill_storage = _make_channel_skill_storage([_make_channel_skill(tmp_path, "data-analysis")])
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text=" /data-analysis analyze uploads/foo.csv",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_not_called()
|
||||
assert outbound_received[0].text.startswith("Unknown command: /data-analysis.")
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_slash_skill_respects_custom_agent_skill_whitelist(self, monkeypatch, tmp_path):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
monkeypatch.setattr("app.channels.manager.load_agent_config", lambda name: SimpleNamespace(skills=["frontend-design"]))
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(
|
||||
bus=bus,
|
||||
store=store,
|
||||
default_session={"assistant_id": "analyst-agent"},
|
||||
)
|
||||
manager._skill_storage = _make_channel_skill_storage([_make_channel_skill(tmp_path, "data-analysis")])
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text="/data-analysis analyze uploads/foo.csv",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_not_called()
|
||||
assert outbound_received[0].text == "Skill `/data-analysis` is not available for this agent."
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_slash_skill_reports_disabled_skill(self, tmp_path):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
manager._skill_storage = _make_channel_skill_storage([_make_channel_skill(tmp_path, "data-analysis", enabled=False)])
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text="/data-analysis analyze uploads/foo.csv",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_not_called()
|
||||
assert outbound_received[0].text == "Skill `/data-analysis` is installed but disabled. Enable it before using slash activation."
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_uninstalled_slash_skill_stays_unknown_command(self, tmp_path):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
manager._skill_storage = _make_channel_skill_storage([_make_channel_skill(tmp_path, "frontend-design")])
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text="/data-analysis analyze uploads/foo.csv",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_not_called()
|
||||
assert outbound_received[0].text.startswith("Unknown command: /data-analysis.")
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_slash_skill_resolution_error_is_reported(self, monkeypatch):
|
||||
from app.channels.manager import ChannelManager, SlashSkillCommandResolutionError
|
||||
|
||||
def fail_resolution(text, available_skills=None, storage=None):
|
||||
raise SlashSkillCommandResolutionError("Failed to resolve slash skill command. Please check the skill configuration.")
|
||||
|
||||
monkeypatch.setattr("app.channels.manager._resolve_slash_skill_command", fail_resolution)
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
store = ChannelStore(path=Path(tempfile.mkdtemp()) / "store.json")
|
||||
manager = ChannelManager(bus=bus, store=store)
|
||||
store.set_thread_id("test", "chat1", "base-thread")
|
||||
store.set_thread_id("test", "chat1", "topic-thread", topic_id="topic-1")
|
||||
|
||||
mock_client = _make_mock_langgraph_client()
|
||||
manager._client = mock_client
|
||||
|
||||
outbound_received = []
|
||||
|
||||
async def capture_outbound(msg):
|
||||
outbound_received.append(msg)
|
||||
|
||||
bus.subscribe_outbound(capture_outbound)
|
||||
await manager.start()
|
||||
|
||||
inbound = InboundMessage(
|
||||
channel_name="test",
|
||||
chat_id="chat1",
|
||||
user_id="user1",
|
||||
text="/data-analysis analyze uploads/foo.csv",
|
||||
msg_type=InboundMessageType.COMMAND,
|
||||
topic_id="topic-1",
|
||||
)
|
||||
await bus.publish_inbound(inbound)
|
||||
await _wait_for(lambda: len(outbound_received) >= 1)
|
||||
await manager.stop()
|
||||
|
||||
mock_client.runs.wait.assert_not_called()
|
||||
assert outbound_received[0].text == "Failed to resolve slash skill command. Please check the skill configuration."
|
||||
assert outbound_received[0].thread_id == "topic-thread"
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_handle_command_new(self):
|
||||
from app.channels.manager import ChannelManager
|
||||
|
||||
@@ -2440,6 +2966,36 @@ class TestWeComChannel:
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_publish_ws_inbound_treats_slash_prefixed_paths_as_chat(self, monkeypatch):
|
||||
from app.channels.wecom import WeComChannel
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
bus.publish_inbound = AsyncMock()
|
||||
channel = WeComChannel(bus, config={})
|
||||
channel._ws_client = SimpleNamespace(reply_stream=AsyncMock())
|
||||
|
||||
monkeypatch.setitem(
|
||||
__import__("sys").modules,
|
||||
"aibot",
|
||||
SimpleNamespace(generate_req_id=lambda prefix: "stream-1"),
|
||||
)
|
||||
|
||||
frame = {
|
||||
"body": {
|
||||
"msgid": "msg-1",
|
||||
"from": {"userid": "user-1"},
|
||||
}
|
||||
}
|
||||
|
||||
await channel._publish_ws_inbound(frame, "/mnt/user-data/uploads/report.pdf")
|
||||
|
||||
inbound = bus.publish_inbound.await_args.args[0]
|
||||
assert inbound.text == "/mnt/user-data/uploads/report.pdf"
|
||||
assert inbound.msg_type == InboundMessageType.CHAT
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_on_outbound_sends_attachment_before_clearing_context(self, tmp_path):
|
||||
from app.channels.wecom import WeComChannel
|
||||
|
||||
@@ -2788,6 +3344,219 @@ class TestSlackAllowedUsers:
|
||||
assert inbound.chat_id == "C123"
|
||||
assert inbound.text == "hello from slack"
|
||||
|
||||
def test_app_mention_strips_leading_bot_mention_before_command_detection(self):
|
||||
from app.channels.slack import SlackChannel
|
||||
|
||||
bus = MessageBus()
|
||||
bus.publish_inbound = AsyncMock()
|
||||
channel = SlackChannel(bus=bus, config={"bot_user_id": "UBOT"})
|
||||
channel._loop = MagicMock()
|
||||
channel._loop.is_running.return_value = True
|
||||
channel._add_reaction = MagicMock()
|
||||
channel._send_running_reply = MagicMock()
|
||||
|
||||
event = {
|
||||
"type": "app_mention",
|
||||
"user": "U123456",
|
||||
"text": "<@UBOT> /help",
|
||||
"channel": "C123",
|
||||
"ts": "1710000000.000100",
|
||||
}
|
||||
|
||||
with patch(
|
||||
"app.channels.slack.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=self._submit_coro,
|
||||
):
|
||||
channel._handle_message_event(event)
|
||||
|
||||
inbound = bus.publish_inbound.call_args.args[0]
|
||||
assert inbound.text == "/help"
|
||||
assert inbound.msg_type == InboundMessageType.COMMAND
|
||||
|
||||
def test_app_mention_strips_labelled_leading_bot_mention(self):
|
||||
from app.channels.slack import SlackChannel
|
||||
|
||||
bus = MessageBus()
|
||||
bus.publish_inbound = AsyncMock()
|
||||
channel = SlackChannel(bus=bus, config={"bot_user_id": "UBOT"})
|
||||
channel._loop = MagicMock()
|
||||
channel._loop.is_running.return_value = True
|
||||
channel._add_reaction = MagicMock()
|
||||
channel._send_running_reply = MagicMock()
|
||||
|
||||
event = {
|
||||
"type": "app_mention",
|
||||
"user": "U123456",
|
||||
"text": "<@UBOT|deerflow> /help",
|
||||
"channel": "C123",
|
||||
"ts": "1710000000.000100",
|
||||
}
|
||||
|
||||
with patch(
|
||||
"app.channels.slack.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=self._submit_coro,
|
||||
):
|
||||
channel._handle_message_event(event)
|
||||
|
||||
inbound = bus.publish_inbound.call_args.args[0]
|
||||
assert inbound.text == "/help"
|
||||
assert inbound.msg_type == InboundMessageType.COMMAND
|
||||
|
||||
def test_app_mention_strips_leading_bot_mention_before_slash_skill(self):
|
||||
from app.channels.slack import SlackChannel
|
||||
|
||||
bus = MessageBus()
|
||||
bus.publish_inbound = AsyncMock()
|
||||
channel = SlackChannel(bus=bus, config={"bot_user_id": "UBOT"})
|
||||
channel._loop = MagicMock()
|
||||
channel._loop.is_running.return_value = True
|
||||
channel._add_reaction = MagicMock()
|
||||
channel._send_running_reply = MagicMock()
|
||||
|
||||
event = {
|
||||
"type": "app_mention",
|
||||
"user": "U123456",
|
||||
"text": "<@UBOT> /data-analysis analyze uploads/foo.csv",
|
||||
"channel": "C123",
|
||||
"ts": "1710000000.000100",
|
||||
}
|
||||
|
||||
with patch(
|
||||
"app.channels.slack.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=self._submit_coro,
|
||||
):
|
||||
channel._handle_message_event(event)
|
||||
|
||||
inbound = bus.publish_inbound.call_args.args[0]
|
||||
assert inbound.text == "/data-analysis analyze uploads/foo.csv"
|
||||
assert inbound.msg_type == InboundMessageType.CHAT
|
||||
|
||||
def test_app_mention_preserves_following_user_mention(self):
|
||||
from app.channels.slack import SlackChannel
|
||||
|
||||
bus = MessageBus()
|
||||
bus.publish_inbound = AsyncMock()
|
||||
channel = SlackChannel(bus=bus, config={"bot_user_id": "UBOT"})
|
||||
channel._loop = MagicMock()
|
||||
channel._loop.is_running.return_value = True
|
||||
channel._add_reaction = MagicMock()
|
||||
channel._send_running_reply = MagicMock()
|
||||
|
||||
event = {
|
||||
"type": "app_mention",
|
||||
"user": "U123456",
|
||||
"text": "<@UBOT> <@UASSIGNEE> please review this",
|
||||
"channel": "C123",
|
||||
"ts": "1710000000.000100",
|
||||
}
|
||||
|
||||
with patch(
|
||||
"app.channels.slack.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=self._submit_coro,
|
||||
):
|
||||
channel._handle_message_event(event)
|
||||
|
||||
inbound = bus.publish_inbound.call_args.args[0]
|
||||
assert inbound.text == "<@UASSIGNEE> please review this"
|
||||
assert inbound.msg_type == InboundMessageType.CHAT
|
||||
|
||||
def test_app_mention_preserves_leading_non_bot_mention_when_bot_id_known(self):
|
||||
from app.channels.slack import SlackChannel
|
||||
|
||||
bus = MessageBus()
|
||||
bus.publish_inbound = AsyncMock()
|
||||
channel = SlackChannel(bus=bus, config={"bot_user_id": "UBOT"})
|
||||
channel._loop = MagicMock()
|
||||
channel._loop.is_running.return_value = True
|
||||
channel._add_reaction = MagicMock()
|
||||
channel._send_running_reply = MagicMock()
|
||||
|
||||
event = {
|
||||
"type": "app_mention",
|
||||
"user": "U123456",
|
||||
"text": "<@UASSIGNEE> <@UBOT> please review this",
|
||||
"channel": "C123",
|
||||
"ts": "1710000000.000100",
|
||||
}
|
||||
|
||||
with patch(
|
||||
"app.channels.slack.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=self._submit_coro,
|
||||
):
|
||||
channel._handle_message_event(event)
|
||||
|
||||
inbound = bus.publish_inbound.call_args.args[0]
|
||||
assert inbound.text == "<@UASSIGNEE> <@UBOT> please review this"
|
||||
assert inbound.msg_type == InboundMessageType.CHAT
|
||||
|
||||
def test_app_mention_preserves_leading_non_bot_mention_when_bot_id_unknown(self):
|
||||
from app.channels.slack import SlackChannel
|
||||
|
||||
bus = MessageBus()
|
||||
bus.publish_inbound = AsyncMock()
|
||||
channel = SlackChannel(bus=bus, config={})
|
||||
channel._loop = MagicMock()
|
||||
channel._loop.is_running.return_value = True
|
||||
channel._add_reaction = MagicMock()
|
||||
channel._send_running_reply = MagicMock()
|
||||
|
||||
event = {
|
||||
"type": "app_mention",
|
||||
"user": "U123456",
|
||||
"text": "<@UASSIGNEE> /help <@UBOT>",
|
||||
"channel": "C123",
|
||||
"ts": "1710000000.000100",
|
||||
}
|
||||
|
||||
with patch(
|
||||
"app.channels.slack.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=self._submit_coro,
|
||||
):
|
||||
channel._handle_message_event(event)
|
||||
|
||||
inbound = bus.publish_inbound.call_args.args[0]
|
||||
assert inbound.text == "<@UASSIGNEE> /help <@UBOT>"
|
||||
assert inbound.msg_type == InboundMessageType.CHAT
|
||||
|
||||
def test_socket_event_resolves_bot_user_id_before_app_mention_command_detection(self):
|
||||
from app.channels.slack import SlackChannel
|
||||
|
||||
bus = MessageBus()
|
||||
bus.publish_inbound = AsyncMock()
|
||||
channel = SlackChannel(bus=bus, config={})
|
||||
channel._SocketModeResponse = lambda envelope_id: SimpleNamespace(envelope_id=envelope_id)
|
||||
channel._loop = MagicMock()
|
||||
channel._loop.is_running.return_value = True
|
||||
channel._add_reaction = MagicMock()
|
||||
channel._send_running_reply = MagicMock()
|
||||
|
||||
client = SimpleNamespace(send_socket_mode_response=MagicMock())
|
||||
req = SimpleNamespace(
|
||||
envelope_id="env-1",
|
||||
type="events_api",
|
||||
payload={
|
||||
"authorizations": [{"user_id": "UBOT"}],
|
||||
"event": {
|
||||
"type": "app_mention",
|
||||
"user": "U123456",
|
||||
"text": "<@UBOT> /help",
|
||||
"channel": "C123",
|
||||
"ts": "1710000000.000100",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"app.channels.slack.asyncio.run_coroutine_threadsafe",
|
||||
side_effect=self._submit_coro,
|
||||
):
|
||||
channel._on_socket_event(client, req)
|
||||
|
||||
inbound = bus.publish_inbound.call_args.args[0]
|
||||
assert channel._bot_user_id == "UBOT"
|
||||
assert inbound.text == "/help"
|
||||
assert inbound.msg_type == InboundMessageType.COMMAND
|
||||
|
||||
def test_scalar_allowed_users_warns_and_matches_stringified_event_user_id(self, caplog):
|
||||
from app.channels.slack import SlackChannel
|
||||
|
||||
@@ -2861,6 +3630,86 @@ class TestSlackAllowedUsers:
|
||||
|
||||
|
||||
class TestTelegramSendRetry:
|
||||
def test_start_registers_known_channel_commands(self, monkeypatch):
|
||||
import sys
|
||||
from types import ModuleType
|
||||
|
||||
from app.channels.commands import KNOWN_CHANNEL_COMMANDS
|
||||
from app.channels.telegram import TelegramChannel
|
||||
|
||||
class FakeFilter:
|
||||
def __init__(self, expr: str):
|
||||
self.expr = expr
|
||||
|
||||
def __and__(self, other):
|
||||
return FakeFilter(f"{self.expr}&{other.expr}")
|
||||
|
||||
def __invert__(self):
|
||||
return FakeFilter(f"~{self.expr}")
|
||||
|
||||
class FakeApplication:
|
||||
def __init__(self):
|
||||
self.handlers = []
|
||||
|
||||
def add_handler(self, handler):
|
||||
self.handlers.append(handler)
|
||||
|
||||
fake_app = FakeApplication()
|
||||
|
||||
class FakeApplicationBuilder:
|
||||
def token(self, token):
|
||||
assert token == "test-token"
|
||||
return self
|
||||
|
||||
def build(self):
|
||||
return fake_app
|
||||
|
||||
def fake_command_handler(command, callback):
|
||||
return SimpleNamespace(kind="command", command=command, callback=callback)
|
||||
|
||||
def fake_message_handler(filter_expr, callback):
|
||||
return SimpleNamespace(kind="message", filter_expr=filter_expr, callback=callback)
|
||||
|
||||
telegram_mod = ModuleType("telegram")
|
||||
telegram_ext_mod = ModuleType("telegram.ext")
|
||||
telegram_ext_mod.ApplicationBuilder = FakeApplicationBuilder
|
||||
telegram_ext_mod.CommandHandler = fake_command_handler
|
||||
telegram_ext_mod.MessageHandler = fake_message_handler
|
||||
telegram_ext_mod.filters = SimpleNamespace(TEXT=FakeFilter("TEXT"), COMMAND=FakeFilter("COMMAND"))
|
||||
telegram_mod.ext = telegram_ext_mod
|
||||
monkeypatch.setitem(sys.modules, "telegram", telegram_mod)
|
||||
monkeypatch.setitem(sys.modules, "telegram.ext", telegram_ext_mod)
|
||||
|
||||
class FakeThread:
|
||||
def __init__(self, *, target, daemon):
|
||||
self.target = target
|
||||
self.daemon = daemon
|
||||
|
||||
def start(self):
|
||||
return None
|
||||
|
||||
def join(self, timeout=None):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr("app.channels.telegram.threading.Thread", FakeThread)
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
ch = TelegramChannel(bus=bus, config={"bot_token": "test-token"})
|
||||
|
||||
await ch.start()
|
||||
try:
|
||||
registered_commands = {handler.command for handler in fake_app.handlers if handler.kind == "command"}
|
||||
expected_commands = {command.removeprefix("/") for command in KNOWN_CHANNEL_COMMANDS}
|
||||
assert expected_commands <= registered_commands
|
||||
assert "start" in registered_commands
|
||||
message_filters = {handler.filter_expr.expr for handler in fake_app.handlers if handler.kind == "message"}
|
||||
assert {"TEXT&COMMAND", "TEXT&~COMMAND"} <= message_filters
|
||||
finally:
|
||||
await ch.stop()
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_retries_on_failure_then_succeeds(self):
|
||||
from app.channels.telegram import TelegramChannel
|
||||
|
||||
@@ -2984,6 +3833,47 @@ class TestTelegramPrivateChatThread:
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_private_chat_slash_skill_text_routes_as_chat(self):
|
||||
from app.channels.telegram import TelegramChannel
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
ch = TelegramChannel(bus=bus, config={"bot_token": "test-token"})
|
||||
ch._main_loop = asyncio.get_event_loop()
|
||||
|
||||
update = _make_telegram_update("private", message_id=12, text="/data-analysis analyze uploads/foo.csv")
|
||||
await ch._on_text(update, None)
|
||||
|
||||
msg = await asyncio.wait_for(bus.get_inbound(), timeout=2)
|
||||
assert msg.text == "/data-analysis analyze uploads/foo.csv"
|
||||
assert msg.msg_type == InboundMessageType.CHAT
|
||||
assert msg.topic_id is None
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_slash_skill_addressed_to_telegram_bot_strips_username(self):
|
||||
from app.channels.telegram import TelegramChannel
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
ch = TelegramChannel(bus=bus, config={"bot_token": "test-token"})
|
||||
ch._main_loop = asyncio.get_event_loop()
|
||||
|
||||
update = _make_telegram_update(
|
||||
"group",
|
||||
message_id=13,
|
||||
text="/data-analysis@DeerFlowBot analyze uploads/foo.csv",
|
||||
)
|
||||
context = SimpleNamespace(bot=SimpleNamespace(username="DeerFlowBot"))
|
||||
await ch._on_text(update, context)
|
||||
|
||||
msg = await asyncio.wait_for(bus.get_inbound(), timeout=2)
|
||||
assert msg.text == "/data-analysis analyze uploads/foo.csv"
|
||||
assert msg.msg_type == InboundMessageType.CHAT
|
||||
assert msg.topic_id == "13"
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_private_chat_with_reply_still_uses_none_topic(self):
|
||||
from app.channels.telegram import TelegramChannel
|
||||
|
||||
@@ -3099,6 +3989,25 @@ class TestTelegramPrivateChatThread:
|
||||
|
||||
_run(go())
|
||||
|
||||
def test_cmd_generic_strips_addressed_telegram_bot_username(self):
|
||||
from app.channels.telegram import TelegramChannel
|
||||
|
||||
async def go():
|
||||
bus = MessageBus()
|
||||
ch = TelegramChannel(bus=bus, config={"bot_token": "test-token"})
|
||||
ch._main_loop = asyncio.get_event_loop()
|
||||
|
||||
update = _make_telegram_update("group", message_id=33, text="/status@DeerFlowBot")
|
||||
context = SimpleNamespace(bot=SimpleNamespace(username="DeerFlowBot"))
|
||||
await ch._cmd_generic(update, context)
|
||||
|
||||
msg = await asyncio.wait_for(bus.get_inbound(), timeout=2)
|
||||
assert msg.text == "/status"
|
||||
assert msg.topic_id == "33"
|
||||
assert msg.msg_type == InboundMessageType.COMMAND
|
||||
|
||||
_run(go())
|
||||
|
||||
|
||||
class TestTelegramProcessingOrder:
|
||||
"""Ensure 'working on it...' is sent before inbound is published."""
|
||||
|
||||
@@ -2,9 +2,13 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from app.channels.discord import DiscordChannel
|
||||
from app.channels.manager import CHANNEL_CAPABILITIES
|
||||
from app.channels.message_bus import MessageBus
|
||||
from app.channels.message_bus import InboundMessageType, MessageBus
|
||||
from app.channels.service import _CHANNEL_REGISTRY
|
||||
|
||||
|
||||
@@ -21,3 +25,64 @@ def test_discord_channel_init() -> None:
|
||||
channel = DiscordChannel(bus=bus, config={"bot_token": "token"})
|
||||
|
||||
assert channel.name == "discord"
|
||||
|
||||
|
||||
def _make_discord_message(text: str):
|
||||
return SimpleNamespace(
|
||||
id=111,
|
||||
content=text,
|
||||
author=SimpleNamespace(id=123, bot=False, display_name="alice"),
|
||||
guild=SimpleNamespace(id=321),
|
||||
channel=SimpleNamespace(id=456),
|
||||
add_reaction=lambda _emoji: None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discord_bot_mention_slash_skill_routes_as_chat() -> None:
|
||||
bus = MessageBus()
|
||||
channel = DiscordChannel(bus=bus, config={"bot_token": "token"})
|
||||
captured = []
|
||||
channel._running = True
|
||||
channel._client = SimpleNamespace(user=SimpleNamespace(id=999, mention="<@999>"))
|
||||
channel._discord_module = SimpleNamespace(Thread=type("FakeThread", (), {}))
|
||||
channel._publish = captured.append
|
||||
|
||||
async def noop(*_args, **_kwargs):
|
||||
return None
|
||||
|
||||
channel._start_typing = noop
|
||||
channel._add_reaction = noop
|
||||
|
||||
await channel._on_message(_make_discord_message("<@999> /data-analysis analyze uploads/foo.csv"))
|
||||
|
||||
assert len(captured) == 1
|
||||
inbound = captured[0]
|
||||
assert inbound.text == "/data-analysis analyze uploads/foo.csv"
|
||||
assert inbound.msg_type == InboundMessageType.CHAT
|
||||
assert inbound.topic_id == "456"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discord_bot_mention_known_command_routes_as_command() -> None:
|
||||
bus = MessageBus()
|
||||
channel = DiscordChannel(bus=bus, config={"bot_token": "token"})
|
||||
captured = []
|
||||
channel._running = True
|
||||
channel._client = SimpleNamespace(user=SimpleNamespace(id=999, mention="<@999>"))
|
||||
channel._discord_module = SimpleNamespace(Thread=type("FakeThread", (), {}))
|
||||
channel._publish = captured.append
|
||||
|
||||
async def noop(*_args, **_kwargs):
|
||||
return None
|
||||
|
||||
channel._start_typing = noop
|
||||
channel._add_reaction = noop
|
||||
|
||||
await channel._on_message(_make_discord_message("<@999> /help"))
|
||||
|
||||
assert len(captured) == 1
|
||||
inbound = captured[0]
|
||||
assert inbound.text == "/help"
|
||||
assert inbound.msg_type == InboundMessageType.COMMAND
|
||||
assert inbound.topic_id == "456"
|
||||
|
||||
@@ -60,6 +60,17 @@ def test_get_skills_prompt_section_returns_all_when_available_skills_is_none(mon
|
||||
assert "skill2" in result
|
||||
|
||||
|
||||
def test_get_skills_prompt_section_includes_slash_activation_guidance(monkeypatch):
|
||||
skills = [_make_skill("data-analysis")]
|
||||
monkeypatch.setattr("deerflow.agents.lead_agent.prompt._get_enabled_skills", lambda: skills)
|
||||
|
||||
result = get_skills_prompt_section(available_skills={"data-analysis"})
|
||||
|
||||
assert "Explicit Slash Skill Activation" in result
|
||||
assert "The runtime injects the activated skill content" in result
|
||||
assert "do not call `read_file` for that SKILL.md again" in result
|
||||
|
||||
|
||||
def test_get_skills_prompt_section_includes_self_evolution_rules(monkeypatch):
|
||||
skills = [_make_skill("skill1")]
|
||||
monkeypatch.setattr("deerflow.agents.lead_agent.prompt._get_enabled_skills", lambda: skills)
|
||||
|
||||
@@ -0,0 +1,557 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
from langchain.agents.middleware.types import ModelRequest
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
|
||||
from app.channels.commands import KNOWN_CHANNEL_COMMANDS
|
||||
from deerflow.agents.middlewares import skill_activation_middleware as middleware_module
|
||||
from deerflow.agents.middlewares.skill_activation_middleware import SkillActivationMiddleware, is_slash_skill_activation_reminder
|
||||
from deerflow.skills.slash import RESERVED_SLASH_SKILL_NAMES, parse_slash_skill_reference, resolve_slash_skill
|
||||
from deerflow.skills.types import Skill, SkillCategory
|
||||
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY
|
||||
|
||||
|
||||
def _make_skill(tmp_path: Path, name: str, content: str = "skill body") -> Skill:
|
||||
skill_dir = tmp_path / name
|
||||
skill_dir.mkdir()
|
||||
skill_file = skill_dir / "SKILL.md"
|
||||
skill_file.write_text(content, encoding="utf-8")
|
||||
return Skill(
|
||||
name=name,
|
||||
description=f"Description for {name}",
|
||||
license="MIT",
|
||||
skill_dir=skill_dir,
|
||||
skill_file=skill_file,
|
||||
relative_path=Path(name),
|
||||
category=SkillCategory.CUSTOM,
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
|
||||
def _make_storage(tmp_path: Path, skills: list[Skill]):
|
||||
return SimpleNamespace(
|
||||
load_skills=lambda *, enabled_only: [skill for skill in skills if skill.enabled] if enabled_only else skills,
|
||||
get_container_root=lambda: "/mnt/skills",
|
||||
get_skills_root_path=lambda: tmp_path,
|
||||
)
|
||||
|
||||
|
||||
def _make_model_request(messages: list[HumanMessage], *, runtime=None) -> ModelRequest:
|
||||
return ModelRequest(
|
||||
model=object(),
|
||||
messages=messages,
|
||||
state={"messages": list(messages)},
|
||||
runtime=runtime,
|
||||
)
|
||||
|
||||
|
||||
def test_parse_slash_skill_reference_extracts_name_and_remaining_text():
|
||||
parsed = parse_slash_skill_reference("/data-analysis analyze uploads/foo.csv")
|
||||
|
||||
assert parsed is not None
|
||||
assert parsed.name == "data-analysis"
|
||||
assert parsed.remaining_text == "analyze uploads/foo.csv"
|
||||
|
||||
|
||||
def test_parse_slash_skill_reference_accepts_skill_name_without_task():
|
||||
parsed = parse_slash_skill_reference("/data-analysis")
|
||||
|
||||
assert parsed is not None
|
||||
assert parsed.name == "data-analysis"
|
||||
assert parsed.remaining_text == ""
|
||||
|
||||
|
||||
def test_parse_slash_skill_reference_rejects_invalid_names():
|
||||
assert parse_slash_skill_reference("/DataAnalysis run") is None
|
||||
assert parse_slash_skill_reference("/data_analysis run") is None
|
||||
assert parse_slash_skill_reference("please use /data-analysis") is None
|
||||
assert parse_slash_skill_reference(" /data-analysis run") is None
|
||||
assert parse_slash_skill_reference("/data-analysis分析这个文档") is None
|
||||
|
||||
|
||||
def test_resolve_slash_skill_ignores_reserved_control_commands(tmp_path):
|
||||
for command in ["bootstrap", "help", "memory", "models", "new", "status"]:
|
||||
skill = _make_skill(tmp_path, command)
|
||||
|
||||
assert resolve_slash_skill(f"/{command} create an agent", [skill]) is None
|
||||
|
||||
|
||||
def test_reserved_slash_skill_names_match_channel_commands():
|
||||
assert RESERVED_SLASH_SKILL_NAMES == {command.removeprefix("/") for command in KNOWN_CHANNEL_COMMANDS}
|
||||
|
||||
|
||||
def test_resolve_slash_skill_respects_available_skill_whitelist(tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis")
|
||||
|
||||
assert resolve_slash_skill("/data-analysis run", [skill], available_skills=set()) is None
|
||||
|
||||
resolved = resolve_slash_skill("/data-analysis run", [skill], available_skills={"data-analysis"})
|
||||
assert resolved is not None
|
||||
assert resolved.skill.name == "data-analysis"
|
||||
assert resolved.remaining_text == "run"
|
||||
assert resolved.container_file_path == "/mnt/skills/custom/data-analysis/SKILL.md"
|
||||
|
||||
|
||||
def test_resolve_slash_skill_rejects_disabled_skills(tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis")
|
||||
skill.enabled = False
|
||||
|
||||
assert resolve_slash_skill("/data-analysis run", [skill]) is None
|
||||
|
||||
|
||||
def test_skill_activation_middleware_injects_hidden_human_context_for_model_call(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis analyze uploads/foo.csv", id="msg-1")
|
||||
request = _make_model_request([original])
|
||||
captured = {}
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(request, handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert result.content == "ok"
|
||||
activation_msg, user_msg = captured["messages"]
|
||||
assert is_slash_skill_activation_reminder(activation_msg)
|
||||
assert activation_msg.additional_kwargs["hide_from_ui"] is True
|
||||
assert "Use pandas." in activation_msg.content
|
||||
assert "<user_request>\nanalyze uploads/foo.csv\n</user_request>" in activation_msg.content
|
||||
assert user_msg.content == original.content
|
||||
assert request.state["messages"] == [original]
|
||||
|
||||
|
||||
def test_skill_activation_middleware_does_not_duplicate_existing_activation(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis analyze uploads/foo.csv", id="msg-1")
|
||||
first_capture = {}
|
||||
|
||||
def first_handler(model_request: ModelRequest):
|
||||
first_capture["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
first_result = middleware.wrap_model_call(_make_model_request([original]), first_handler)
|
||||
|
||||
assert isinstance(first_result, AIMessage)
|
||||
activation_msg, user_msg = first_capture["messages"]
|
||||
assert is_slash_skill_activation_reminder(activation_msg)
|
||||
|
||||
second_capture = {}
|
||||
|
||||
def second_handler(model_request: ModelRequest):
|
||||
second_capture["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
second_result = middleware.wrap_model_call(_make_model_request([activation_msg, user_msg]), second_handler)
|
||||
|
||||
assert isinstance(second_result, AIMessage)
|
||||
assert second_capture["messages"] == [activation_msg, user_msg]
|
||||
assert sum(is_slash_skill_activation_reminder(message) for message in second_capture["messages"]) == 1
|
||||
|
||||
|
||||
def test_skill_activation_middleware_does_not_duplicate_activation_separated_by_hidden_context(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis analyze uploads/foo.csv", id="msg-1")
|
||||
first_capture = {}
|
||||
|
||||
def first_handler(model_request: ModelRequest):
|
||||
first_capture["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
middleware.wrap_model_call(_make_model_request([original]), first_handler)
|
||||
activation_msg, user_msg = first_capture["messages"]
|
||||
hidden_context = HumanMessage(content="dynamic context", additional_kwargs={"hide_from_ui": True})
|
||||
second_capture = {}
|
||||
|
||||
def second_handler(model_request: ModelRequest):
|
||||
second_capture["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
second_result = middleware.wrap_model_call(_make_model_request([activation_msg, hidden_context, user_msg]), second_handler)
|
||||
|
||||
assert isinstance(second_result, AIMessage)
|
||||
assert second_capture["messages"] == [activation_msg, hidden_context, user_msg]
|
||||
assert sum(is_slash_skill_activation_reminder(message) for message in second_capture["messages"]) == 1
|
||||
|
||||
|
||||
def test_skill_activation_middleware_dedupes_immediately_previous_activation_without_target_id(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
legacy_activation_msg = SkillActivationMiddleware._make_activation_message(
|
||||
HumanMessage(content="/data-analysis analyze uploads/foo.csv"),
|
||||
"existing activation context",
|
||||
)
|
||||
target = HumanMessage(content="/data-analysis analyze uploads/foo.csv", id="msg-1")
|
||||
captured = {}
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([legacy_activation_msg, target]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert captured["messages"] == [legacy_activation_msg, target]
|
||||
assert sum(is_slash_skill_activation_reminder(message) for message in captured["messages"]) == 1
|
||||
|
||||
|
||||
def test_skill_activation_middleware_async_injects_hidden_human_context_for_model_call(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis analyze uploads/foo.csv", id="msg-1")
|
||||
request = _make_model_request([original])
|
||||
captured = {}
|
||||
|
||||
async def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = asyncio.run(middleware.awrap_model_call(request, handler))
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert result.content == "ok"
|
||||
activation_msg, user_msg = captured["messages"]
|
||||
assert is_slash_skill_activation_reminder(activation_msg)
|
||||
assert activation_msg.additional_kwargs["hide_from_ui"] is True
|
||||
assert "Use pandas." in activation_msg.content
|
||||
assert "<user_request>\nanalyze uploads/foo.csv\n</user_request>" in activation_msg.content
|
||||
assert user_msg.content == original.content
|
||||
assert request.state["messages"] == [original]
|
||||
|
||||
|
||||
def test_skill_activation_middleware_uses_fallback_when_task_text_is_empty(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis", id="msg-1")
|
||||
captured = {}
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
activation_msg = captured["messages"][0]
|
||||
assert "No additional task text was provided after the slash skill command." in activation_msg.content
|
||||
|
||||
|
||||
def test_skill_activation_middleware_uses_original_user_content_when_uploads_are_injected(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(
|
||||
content="<uploaded_files>\n- report.pdf\n</uploaded_files>\n\n/data-analysis 分析这个文档",
|
||||
id="msg-1",
|
||||
additional_kwargs={ORIGINAL_USER_CONTENT_KEY: "/data-analysis 分析这个文档"},
|
||||
)
|
||||
captured = {}
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert result.content == "ok"
|
||||
activation_msg, user_msg = captured["messages"]
|
||||
assert is_slash_skill_activation_reminder(activation_msg)
|
||||
assert "Use pandas." in activation_msg.content
|
||||
assert "<user_request>\n分析这个文档\n</user_request>" in activation_msg.content
|
||||
assert user_msg.content == original.content
|
||||
assert user_msg.additional_kwargs[ORIGINAL_USER_CONTENT_KEY] == "/data-analysis 分析这个文档"
|
||||
|
||||
|
||||
def test_skill_activation_middleware_activates_from_list_content(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content=[{"type": "text", "text": "/data-analysis analyze uploads/foo.csv"}], id="msg-1")
|
||||
captured = {}
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
activation_msg, user_msg = captured["messages"]
|
||||
assert is_slash_skill_activation_reminder(activation_msg)
|
||||
assert "<user_request>\nanalyze uploads/foo.csv\n</user_request>" in activation_msg.content
|
||||
assert user_msg.content == original.content
|
||||
|
||||
|
||||
def test_skill_activation_middleware_records_activation_audit_event(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
recorded = []
|
||||
journal = SimpleNamespace(record_middleware=lambda *args, **kwargs: recorded.append((args, kwargs)))
|
||||
runtime = SimpleNamespace(context={"__run_journal": journal})
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis analyze uploads/foo.csv", id="msg-1")
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original], runtime=runtime), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert len(recorded) == 1
|
||||
args, kwargs = recorded[0]
|
||||
assert args == ("skill_activation",)
|
||||
assert kwargs["name"] == "SkillActivationMiddleware"
|
||||
assert kwargs["hook"] == "wrap_model_call"
|
||||
assert kwargs["action"] == "activate"
|
||||
assert kwargs["changes"] == {
|
||||
"skill_name": "data-analysis",
|
||||
"category": "custom",
|
||||
"path": "/mnt/skills/custom/data-analysis/SKILL.md",
|
||||
"content_hash": hashlib.sha256(b"# Data Analysis\nUse pandas.").hexdigest(),
|
||||
}
|
||||
|
||||
|
||||
def test_skill_activation_middleware_async_records_activation_audit_event(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
recorded = []
|
||||
journal = SimpleNamespace(record_middleware=lambda *args, **kwargs: recorded.append((args, kwargs)))
|
||||
runtime = SimpleNamespace(context={"__run_journal": journal})
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis analyze uploads/foo.csv", id="msg-1")
|
||||
|
||||
async def handler(model_request: ModelRequest):
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = asyncio.run(middleware.awrap_model_call(_make_model_request([original], runtime=runtime), handler))
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert len(recorded) == 1
|
||||
args, kwargs = recorded[0]
|
||||
assert args == ("skill_activation",)
|
||||
assert kwargs["hook"] == "awrap_model_call"
|
||||
assert kwargs["changes"]["skill_name"] == "data-analysis"
|
||||
assert kwargs["changes"]["content_hash"] == hashlib.sha256(b"# Data Analysis\nUse pandas.").hexdigest()
|
||||
|
||||
|
||||
def test_skill_activation_middleware_ignores_activation_audit_errors(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
journal = SimpleNamespace(record_middleware=lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("db down")))
|
||||
runtime = SimpleNamespace(context={"__run_journal": journal})
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis analyze uploads/foo.csv", id="msg-1")
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original], runtime=runtime), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert result.content == "ok"
|
||||
|
||||
|
||||
def test_skill_activation_middleware_activates_only_latest_real_user_message(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
old_slash = HumanMessage(content="/data-analysis old request", id="msg-1")
|
||||
latest_user = HumanMessage(content="continue normally", id="msg-2")
|
||||
request = _make_model_request([old_slash, AIMessage(content="done"), latest_user])
|
||||
captured = {}
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(request, handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert captured["messages"] == request.messages
|
||||
assert not any(is_slash_skill_activation_reminder(message) for message in captured["messages"])
|
||||
|
||||
|
||||
def test_skill_activation_middleware_ignores_hidden_and_summary_user_messages(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis", content="# Data Analysis\nUse pandas.")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
real_user = HumanMessage(content="continue normally", id="msg-1")
|
||||
hidden_slash = HumanMessage(content="/data-analysis hidden request", id="msg-2", additional_kwargs={"hide_from_ui": True})
|
||||
summary_slash = HumanMessage(content="/data-analysis summary request", id="msg-3", name="summary")
|
||||
request = _make_model_request([real_user, hidden_slash, summary_slash])
|
||||
captured = {}
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(request, handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert captured["messages"] == request.messages
|
||||
assert not any(is_slash_skill_activation_reminder(message) for message in captured["messages"])
|
||||
|
||||
|
||||
def test_skill_activation_middleware_returns_clear_error_for_disallowed_skill(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware(available_skills={"frontend-design"})
|
||||
original = HumanMessage(content="/data-analysis run")
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
raise AssertionError("handler should not be called for invalid slash skills")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert "not available for this agent" in result.content
|
||||
|
||||
|
||||
def test_skill_activation_middleware_returns_clear_error_for_missing_skill(monkeypatch, tmp_path):
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, []))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis run")
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
raise AssertionError("handler should not be called for missing slash skills")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert "not installed" in result.content
|
||||
|
||||
|
||||
def test_skill_activation_middleware_returns_clear_error_for_disabled_skill(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis")
|
||||
skill.enabled = False
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis run")
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
raise AssertionError("handler should not be called for disabled slash skills")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert "installed but disabled" in result.content
|
||||
|
||||
|
||||
def test_skill_activation_middleware_escapes_activation_content(monkeypatch, tmp_path):
|
||||
skill = _make_skill(
|
||||
tmp_path,
|
||||
"data-analysis",
|
||||
content="# Data Analysis\nUse <xml> & avoid </skill> collisions.\n----- END SKILL.md -----",
|
||||
)
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
original = HumanMessage(content="/data-analysis analyze </user_request>")
|
||||
captured = {}
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
captured["messages"] = model_request.messages
|
||||
return AIMessage(content="ok")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([original]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
activation_msg = captured["messages"][0]
|
||||
assert '<skill_content encoding="xml-escaped">' in activation_msg.content
|
||||
assert "analyze </user_request>" in activation_msg.content
|
||||
assert "Use <xml> & avoid </skill> collisions." in activation_msg.content
|
||||
assert "----- BEGIN SKILL.md -----" not in activation_msg.content
|
||||
|
||||
|
||||
def test_skill_activation_middleware_rejects_skill_file_outside_skills_root(monkeypatch, tmp_path):
|
||||
skills_root = tmp_path / "skills"
|
||||
skill_dir = skills_root / "custom" / "data-analysis"
|
||||
skill_dir.mkdir(parents=True)
|
||||
outside_dir = tmp_path / "outside"
|
||||
outside_dir.mkdir()
|
||||
outside_file = outside_dir / "SKILL.md"
|
||||
outside_file.write_text("# Leaked\nDo not read me.", encoding="utf-8")
|
||||
(skill_dir / "SKILL.md").symlink_to(outside_file)
|
||||
skill = Skill(
|
||||
name="data-analysis",
|
||||
description="Description for data-analysis",
|
||||
license="MIT",
|
||||
skill_dir=skill_dir,
|
||||
skill_file=skill_dir / "SKILL.md",
|
||||
relative_path=Path("data-analysis"),
|
||||
category=SkillCategory.CUSTOM,
|
||||
enabled=True,
|
||||
)
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(skills_root, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
raise AssertionError("handler should not be called when SKILL.md fails safety checks")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([HumanMessage(content="/data-analysis run")]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert "could not be loaded safely" in result.content
|
||||
|
||||
|
||||
def test_skill_activation_middleware_reports_missing_skill_file_safely(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis")
|
||||
skill.skill_file.unlink()
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
raise AssertionError("handler should not be called when SKILL.md is missing")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([HumanMessage(content="/data-analysis run")]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert "could not be loaded safely" in result.content
|
||||
|
||||
|
||||
def test_skill_activation_middleware_reports_invalid_utf8_skill_file_safely(monkeypatch, tmp_path):
|
||||
skill = _make_skill(tmp_path, "data-analysis")
|
||||
skill.skill_file.write_bytes(b"\xff\xfe\x00")
|
||||
monkeypatch.setattr(middleware_module, "get_or_new_skill_storage", lambda **kwargs: _make_storage(tmp_path, [skill]))
|
||||
|
||||
middleware = SkillActivationMiddleware()
|
||||
|
||||
def handler(model_request: ModelRequest):
|
||||
raise AssertionError("handler should not be called when SKILL.md is not valid UTF-8")
|
||||
|
||||
result = middleware.wrap_model_call(_make_model_request([HumanMessage(content="/data-analysis run")]), handler)
|
||||
|
||||
assert isinstance(result, AIMessage)
|
||||
assert "could not be loaded safely" in result.content
|
||||
@@ -14,6 +14,7 @@ from langchain_core.messages import AIMessage, HumanMessage
|
||||
|
||||
from deerflow.agents.middlewares.uploads_middleware import UploadsMiddleware
|
||||
from deerflow.config.paths import Paths
|
||||
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY
|
||||
|
||||
THREAD_ID = "thread-abc123"
|
||||
|
||||
@@ -263,6 +264,22 @@ class TestBeforeAgent:
|
||||
assert "<uploaded_files>" in combined_text
|
||||
assert "analyse this" in combined_text
|
||||
|
||||
def test_list_content_preserves_original_slash_skill_text(self, tmp_path):
|
||||
mw = _middleware(tmp_path)
|
||||
uploads_dir = _uploads_dir(tmp_path)
|
||||
(uploads_dir / "data.csv").write_bytes(b"a,b")
|
||||
|
||||
msg = _human(
|
||||
[{"type": "text", "text": "/data-analysis analyze data.csv"}],
|
||||
files=[{"filename": "data.csv", "size": 3, "path": "/mnt/user-data/uploads/data.csv"}],
|
||||
)
|
||||
result = mw.before_agent(self._state(msg), _runtime())
|
||||
|
||||
assert result is not None
|
||||
updated_msg = result["messages"][-1]
|
||||
assert isinstance(updated_msg.content, list)
|
||||
assert updated_msg.additional_kwargs[ORIGINAL_USER_CONTENT_KEY] == "/data-analysis analyze data.csv"
|
||||
|
||||
def test_preserves_additional_kwargs_on_updated_message(self, tmp_path):
|
||||
mw = _middleware(tmp_path)
|
||||
uploads_dir = _uploads_dir(tmp_path)
|
||||
@@ -278,6 +295,37 @@ class TestBeforeAgent:
|
||||
assert updated_kwargs.get("files") == files_meta
|
||||
assert updated_kwargs.get("element") == "task"
|
||||
|
||||
def test_preserves_original_user_content_before_upload_context(self, tmp_path):
|
||||
mw = _middleware(tmp_path)
|
||||
uploads_dir = _uploads_dir(tmp_path)
|
||||
(uploads_dir / "report.pdf").write_bytes(b"pdf")
|
||||
|
||||
msg = _human(
|
||||
"/data-analysis 分析这个文档",
|
||||
files=[{"filename": "report.pdf", "size": 3, "path": "/mnt/user-data/uploads/report.pdf"}],
|
||||
)
|
||||
result = mw.before_agent(self._state(msg), _runtime())
|
||||
|
||||
assert result is not None
|
||||
updated_msg = result["messages"][-1]
|
||||
assert updated_msg.content.startswith("<uploaded_files>")
|
||||
assert updated_msg.additional_kwargs[ORIGINAL_USER_CONTENT_KEY] == "/data-analysis 分析这个文档"
|
||||
|
||||
def test_preserves_existing_original_user_content_marker(self, tmp_path):
|
||||
mw = _middleware(tmp_path)
|
||||
uploads_dir = _uploads_dir(tmp_path)
|
||||
(uploads_dir / "report.pdf").write_bytes(b"pdf")
|
||||
|
||||
msg = _human(
|
||||
"<uploaded_files>\nold\n</uploaded_files>\n\n/data-analysis run",
|
||||
files=[{"filename": "report.pdf", "size": 3, "path": "/mnt/user-data/uploads/report.pdf"}],
|
||||
**{ORIGINAL_USER_CONTENT_KEY: "/data-analysis run"},
|
||||
)
|
||||
result = mw.before_agent(self._state(msg), _runtime())
|
||||
|
||||
assert result is not None
|
||||
assert result["messages"][-1].additional_kwargs[ORIGINAL_USER_CONTENT_KEY] == "/data-analysis run"
|
||||
|
||||
def test_uploaded_files_returned_in_state_update(self, tmp_path):
|
||||
mw = _middleware(tmp_path)
|
||||
uploads_dir = _uploads_dir(tmp_path)
|
||||
|
||||
Reference in New Issue
Block a user