first commit

This commit is contained in:
2026-07-22 13:48:46 +08:00
commit c87751c3dc
2820 changed files with 726976 additions and 0 deletions
+52
View File
@@ -0,0 +1,52 @@
# Context System Tests
This directory contains tests for the Context system. The tests verify Redis-backed history storage, file-system persistence, message management, context management, serialization, concurrency behavior, and error handling.
## Files
- `test_context_system.py`: complete pytest suite.
- `test_context_basic.py`: basic validation script that does not require pytest.
- `test_context_simple.py`: simplified quick-check script.
- `run_context_tests.py`: helper script for dependency checks, Redis checks, test execution, and reporting.
## Run
Start Redis first:
```bash
redis-server
```
Run the quick validation:
```bash
python test/test_context_simple.py
```
Run the full pytest suite:
```bash
pytest test/test_context_system.py -v
```
Or use the helper runner:
```bash
python test/run_context_tests.py
```
## Coverage
- Redis connection and message storage
- File persistence and restore
- Message retrieval, search, and count limits
- Context creation, lookup, batch save, and cleanup
- Serialization and metadata handling
- Concurrent access safety
- Error handling and automatic summaries
## Notes
- Tests use Redis DB 13-15 to avoid affecting production data.
- File-system tests use temporary directories and clean up after themselves.
- See [TEST_RESULTS.md](./TEST_RESULTS.md) for the historical result summary.
@@ -0,0 +1,40 @@
# Context System Test Results
## Overview
The Context system tests verify that conversation history can be stored in Redis and persisted to the file system. The historical test run recorded here passed all core checks.
## Environment
- Operating system: macOS 24.5.0
- Python: 3.12.10
- Redis: latest available version at test time
- Test database: Redis DB 12
- Test directory: temporary directory with automatic cleanup
## Passed Checks
| Test Item | Status | Notes |
|---|---|---|
| Redis connection | Passed | Redis service was reachable |
| Basic message storage | Passed | Messages were stored and retrieved successfully |
| File persistence | Passed | Data was persisted to the file system |
| Redis data validation | Passed | Redis data format and contents were valid |
| Metadata management | Passed | Metadata updates and retrieval worked correctly |
| Message search | Passed | Keyword search returned expected results |
## Summary
The tests confirmed that the Context system can:
- Store conversation history in Redis.
- Persist data to the file system.
- Provide message management, metadata management, and message search.
- Serialize data in a consistent JSON structure.
## Related Files
- `test_context_simple.py`: simplified quick validation.
- `test_context_basic.py`: basic functional tests.
- `test_context_system.py`: complete pytest suite.
- `run_context_tests.py`: test runner helper.
+1
View File
@@ -0,0 +1 @@
# Test package for CADDesigner.
@@ -0,0 +1,168 @@
#!/usr/bin/env python3
"""
Context system demo script.
Shows practical Context system usage scenarios:
1. Create a conversation context
2. Add messages
3. Persist to file
4. Restore from file
5. Search historical messages
"""
import os
import sys
import asyncio
import tempfile
import shutil
from datetime import datetime
# Add the project root directory to the Python path.
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, project_root)
from context.context_manager import ContextManager
from context.context import RedisFileContextBackend
from context.schemas import Message
async def demo_context_system():
"""Demonstrate the complete Context system functionality."""
print("=" * 60)
print("Context System Feature Demo")
print("=" * 60)
# Create a temporary directory.
demo_dir = tempfile.mkdtemp(prefix="demo_context_")
try:
# 1. Create a context manager.
print("\n1. Creating context manager...")
manager = ContextManager(backend_class=RedisFileContextBackend)
# 2. Create a conversation context.
print("\n2. Creating conversation context...")
context_id = "demo_conversation_001"
context = manager.create_context(
context_id=context_id,
max_history_length=10
)
print(f"✓ Created context: {context_id}")
# 3. Simulate a conversation.
print("\n3. Simulating conversation...")
conversation = [
("user", "Hello, I want to learn about Python programming"),
("assistant", "Hello! Python is a very popular programming language. What features are you interested in?"),
("user", "What are Python's advantages?"),
("assistant", "Python's main advantages include:\n1. Concise and readable syntax\n2. Rich libraries and frameworks\n3. Cross-platform support\n4. Beginner friendliness"),
("user", "I want to learn machine learning. Is Python suitable?"),
("assistant", "Absolutely! Python is very popular in machine learning and has many excellent libraries such as TensorFlow, PyTorch, and scikit-learn."),
("user", "Thanks for the introduction"),
("assistant", "You're welcome! If you have any Python or machine learning questions, feel free to ask.")
]
# Add conversation messages.
for role, content in conversation:
success = await manager.add_message(
context_id=context_id,
role=role,
content=content
)
if success:
print(f"✓ Added {role} message: {content[:30]}...")
else:
print(f"✗ Failed to add {role} message")
# 4. View conversation history.
print("\n4. Viewing conversation history...")
history = manager.get_history(context_id)
print(f"✓ Conversation history contains {len(history)} messages")
for i, message in enumerate(history[-3:], 1): # Show the last 3 messages.
print(f" {i}. {message.role}: {message.content[:50]}...")
# 5. Search historical messages.
print("\n5. Searching historical messages...")
search_results = context.search_messages("Python", limit=5)
print(f"✓ Found {len(search_results)} messages containing 'Python'")
for i, message in enumerate(search_results, 1):
print(f" {i}. {message.role}: {message.content[:50]}...")
# 6. Persist to file.
print("\n6. Persisting to file...")
success = await context.persist()
if success:
print("✓ Successfully persisted to the file system")
# Check file contents.
file_path = context.file_path
if os.path.exists(file_path):
file_size = os.path.getsize(file_path)
print(f"✓ File size: {file_size} bytes")
else:
print("✗ Persistence failed")
# 7. Verify Redis storage.
print("\n7. Verifying Redis storage...")
message_count = context.get_message_count()
metadata = context.get_metadata()
print(f"✓ Redis stores {message_count} messages")
print(f"✓ Metadata contains {len(metadata)} fields")
# 8. Simulate system restart (restore from file).
print("\n8. Simulating system restart...")
# Create a new context manager to simulate a restart.
new_manager = ContextManager(backend_class=RedisFileContextBackend)
new_context = new_manager.get_context(context_id)
if new_context:
restored_history = new_context.retrieve_messages()
print(f"✓ Successfully restored conversation history with {len(restored_history)} messages")
# Verify restored data.
if len(restored_history) == len(history):
print("✓ Data integrity verification passed")
else:
print("✗ Data integrity verification failed")
else:
print("✗ Failed to restore conversation history")
# 9. Display system statistics.
print("\n9. System statistics...")
contexts = manager.list_contexts()
print(f"✓ Current system has {len(contexts)} contexts")
for ctx_info in contexts:
print(f" - {ctx_info['context_id']}: {ctx_info.get('total_messages', 0)} messages")
print("\n" + "=" * 60)
print("Demo complete!")
print("=" * 60)
print("\n🎉 Context system feature verification succeeded:")
print(" ✓ Message storage and retrieval work correctly")
print(" ✓ File system persistence works correctly")
print(" ✓ Message search works correctly")
print(" ✓ Data recovery after system restart works correctly")
print(" ✓ Redis storage works correctly")
print(f"\n📁 Demo file location: {demo_dir}")
print("💡 Inspect the generated JSON file to understand the data format")
except Exception as e:
print(f"❌ Error during demo: {e}")
import traceback
traceback.print_exc()
finally:
# Clean up demo files.
if os.path.exists(demo_dir):
shutil.rmtree(demo_dir)
if __name__ == "__main__":
asyncio.run(demo_context_system())
@@ -0,0 +1,238 @@
#!/usr/bin/env python3
"""
Context system test runner script.
Features:
1. Check whether the Redis service is available
2. Run all Context system tests
3. Generate a test report
"""
import os
import sys
import subprocess
import time
import redis
from typing import Optional, Tuple
# Add the project root directory to the Python path.
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, project_root)
def check_redis_connection(host: str = "localhost", port: int = 6379, timeout: int = 5) -> Tuple[bool, str]:
"""
Check whether the Redis connection is available.
Args:
host: Redis host address
port: Redis port
timeout: Connection timeout in seconds
Returns:
(whether the connection succeeded, error message)
"""
try:
r = redis.Redis(host=host, port=port, socket_connect_timeout=timeout, decode_responses=True)
r.ping()
r.close()
return True, "Redis connection is healthy"
except redis.ConnectionError as e:
return False, f"Redis connection failed: {e}"
except Exception as e:
return False, f"Redis check exception: {e}"
def start_redis_server() -> bool:
"""
Try to start the Redis server.
Returns:
Whether startup succeeded.
"""
try:
# Check whether Redis is already running.
if check_redis_connection()[0]:
print("✓ Redis service is already running")
return True
# Try to start Redis.
print("Trying to start the Redis service...")
# Try to start Redis on macOS.
if sys.platform == "darwin":
try:
# Start Redis with brew.
subprocess.run(["brew", "services", "start", "redis"],
check=True, capture_output=True, timeout=10)
time.sleep(2) # Wait for the service to start.
if check_redis_connection()[0]:
print("✓ Redis service started successfully")
return True
except subprocess.CalledProcessError:
pass
try:
# Start Redis directly.
subprocess.run(["redis-server", "--daemonize", "yes"],
check=True, capture_output=True, timeout=10)
time.sleep(2)
if check_redis_connection()[0]:
print("✓ Redis service started successfully")
return True
except subprocess.CalledProcessError:
pass
# Try to start Redis on Linux.
elif sys.platform.startswith("linux"):
try:
subprocess.run(["sudo", "systemctl", "start", "redis"],
check=True, capture_output=True, timeout=10)
time.sleep(2)
if check_redis_connection()[0]:
print("✓ Redis service started successfully")
return True
except subprocess.CalledProcessError:
pass
print("⚠️ Unable to start Redis automatically; please start it manually")
return False
except Exception as e:
print(f"⚠️ Error while starting Redis service: {e}")
return False
def install_test_dependencies() -> bool:
"""
Install test dependencies.
Returns:
Whether installation succeeded.
"""
try:
print("Checking test dependencies...")
# Check pytest.
try:
import pytest
print("✓ pytest is installed")
except ImportError:
print("Installing pytest...")
subprocess.run([sys.executable, "-m", "pip", "install", "pytest", "pytest-asyncio"],
check=True)
print("✓ pytest installation complete")
# Check redis.
try:
import redis
print("✓ redis is installed")
except ImportError:
print("Installing redis...")
subprocess.run([sys.executable, "-m", "pip", "install", "redis"],
check=True)
print("✓ redis installation complete")
return True
except Exception as e:
print(f"⚠️ Error while installing test dependencies: {e}")
return False
def run_tests() -> bool:
"""
Run Context system tests.
Returns:
Whether all tests passed.
"""
try:
print("\n" + "="*60)
print("Starting Context system tests")
print("="*60)
# Get the test file path.
test_file = os.path.join(project_root, "test", "test_context_system.py")
if not os.path.exists(test_file):
print(f"❌ Test file does not exist: {test_file}")
return False
# Run tests.
cmd = [
sys.executable, "-m", "pytest",
test_file,
"-v", # Verbose output.
"-s", # Show print output.
"--tb=short", # Short traceback.
"--color=yes" # Colored output.
]
print(f"Executing command: {' '.join(cmd)}")
print("-" * 60)
result = subprocess.run(cmd, cwd=project_root)
print("-" * 60)
if result.returncode == 0:
app_log("✅ All tests passed!")
return True
else:
print("❌ Some tests failed")
return False
except Exception as e:
print(f"❌ Error while running tests: {e}")
return False
def main():
"""Main function."""
print("Context System Test Runner")
print("=" * 40)
# 1. Install dependencies.
if not install_test_dependencies():
print("❌ Dependency installation failed; exiting")
return 1
# 2. Check Redis connection.
print("\nChecking Redis service...")
redis_ok, redis_msg = check_redis_connection()
if not redis_ok:
print(f"⚠️ {redis_msg}")
print("Trying to start Redis service...")
if not start_redis_server():
print("❌ Unable to start Redis service")
print("Please start Redis manually and rerun the tests")
print("Example startup commands:")
print(" macOS: brew services start redis")
print(" Linux: sudo systemctl start redis")
print(" Windows: redis-server")
return 1
else:
print(f"✓ {redis_msg}")
# 3. Run tests.
success = run_tests()
if success:
print("\n🎉 Context system tests complete!")
print("The test results verify:")
print(" ✓ Redis storage works correctly")
print(" ✓ File system persistence works correctly")
print(" ✓ Message management works correctly")
print(" ✓ Context manager works correctly")
print(" ✓ Data serialization works correctly")
print(" ✓ Concurrent access is safe")
return 0
else:
print("\n❌ Context system tests failed")
return 1
if __name__ == "__main__":
exit(main())
@@ -0,0 +1,55 @@
import os
import sys
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
from agent.CADAgent import CADAgent
from tools import create_builtin_file_tools
from tools.code_tools import create_codegen_subagent_tools
def _tool_name(tool):
if hasattr(tool, "_tool"):
return tool._tool.name
return getattr(tool, "name", getattr(tool, "__name__", None))
def test_create_builtin_file_tools_exposes_simplellmfunc_file_tool_names(tmp_path):
tools = create_builtin_file_tools(tmp_path)
assert [tool.name for tool in tools] == [
"read_file",
"grep",
"sed",
"echo_into",
]
def test_codegen_subagent_tools_include_command_and_builtin_file_tools(tmp_path):
tools = create_codegen_subagent_tools(tmp_path)
assert [_tool_name(tool) for tool in tools] == [
"execute_command",
"sketch_pad_operations",
"read_file",
"grep",
"sed",
"echo_into",
]
def test_main_cad_agent_toolkit_exposes_codegen_specialist_not_low_level_file_tools():
agent = object.__new__(CADAgent)
toolkit = CADAgent.get_toolkit(agent)
names = [_tool_name(tool) for tool in toolkit]
assert "cad_code_generator" in names
assert "execute_command" in names
assert "read_file" not in names
assert "grep" not in names
assert "sed" not in names
assert "echo_into" not in names
@@ -0,0 +1,19 @@
from agent.CADAgent import CADAgent
def test_cadagent_prompt_includes_artifact_tags() -> None:
prompt = CADAgent.chat_impl.__doc__ or ""
assert "<|code_file|>" in prompt
assert "<|output_file|>" in prompt
assert "Repeat the final `<|code_file|>` and `<|output_file|>` tags" in prompt
def test_cadagent_prompt_prefers_step_for_visual_feedback() -> None:
prompt = CADAgent.chat_impl.__doc__ or ""
assert (
"prefer passing the exported `.step`/`.stp` file into `get_visual_feedback`"
in prompt
)
assert "visual feedback should use the STEP/STP file" in prompt
@@ -0,0 +1,200 @@
import os
import sys
from datetime import datetime, timezone
from typing import Any, cast
import pytest
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
WORKSPACE_ROOT = os.path.join(PROJECT_ROOT, "workspace")
from SimpleLLMFunc.hooks.events import (
ReActEventType,
ToolCallEndEvent,
ToolCallErrorEvent,
ToolCallStartEvent,
)
from SimpleLLMFunc.hooks.stream import EventYield, ResponseYield
import tools.code_tools as code_tools_module
from tools.code_tools import cad_code_generator
class _FakeEmitter:
def __init__(self):
self.events: list[tuple[str, dict]] = []
async def emit(self, event_name: str, data):
self.events.append((event_name, data))
@pytest.mark.asyncio
async def test_cad_code_generator_bridges_nested_specialist_events(monkeypatch):
monkeypatch.chdir(WORKSPACE_ROOT)
async def fake_specialist(**kwargs):
assert "Current working directory:" in kwargs["message"]
assert "Skill root: use the preferred skill root below." in kwargs["message"]
assert "validation_command: uv run python part/model.py" in kwargs["message"]
assert "references/docs/api/README.md" in kwargs["message"]
yield EventYield(
event=ToolCallStartEvent(
event_type=ReActEventType.TOOL_CALL_START,
timestamp=datetime.now(timezone.utc),
trace_id="trace-1",
func_name="cad_code_generator_specialist",
iteration=0,
tool_name="read_file",
tool_call_id="nested-1",
arguments={"file_path": "part/model.py"},
tool_call=cast(
Any,
{
"id": "nested-1",
"type": "function",
"function": {"name": "read_file", "arguments": "{}"},
},
),
)
)
yield EventYield(
event=ToolCallEndEvent(
event_type=ReActEventType.TOOL_CALL_END,
timestamp=datetime.now(timezone.utc),
trace_id="trace-1",
func_name="cad_code_generator_specialist",
iteration=0,
tool_name="read_file",
tool_call_id="nested-1",
arguments={"file_path": "part/model.py"},
result="print('old')",
execution_time=0.05,
success=True,
)
)
yield EventYield(
event=ToolCallErrorEvent(
event_type=ReActEventType.TOOL_CALL_ERROR,
timestamp=datetime.now(timezone.utc),
trace_id="trace-1",
func_name="cad_code_generator_specialist",
iteration=0,
tool_name="execute_command",
tool_call_id="nested-2",
arguments={"command": "python model.py"},
error=RuntimeError("boom"),
error_message="boom",
error_type="RuntimeError",
execution_time=0.12,
)
)
yield ResponseYield(response="STATUS: SUCCESS\nSUMMARY: ok", messages=[])
monkeypatch.setattr(
code_tools_module,
"cad_code_generator_specialist",
fake_specialist,
)
monkeypatch.setattr(
code_tools_module, "_read_latest_code", lambda path: "print('ok')\n"
)
emitter = _FakeEmitter()
result = await cad_code_generator(
task="Create a cube as a new file. This is a create-new-file task.",
target_file_path="part/model.py",
event_emitter=emitter,
)
event_names = [event_name for event_name, _ in emitter.events]
assert event_names == [
"subagent_status",
"subagent_tool_start",
"subagent_tool_end",
"subagent_tool_error",
"subagent_response",
"subagent_status",
]
first_status_payload = emitter.events[0][1]
assert first_status_payload["subagent_label"] == "CAD Code Specialist"
assert first_status_payload["validation_command"].startswith(
"uv run python part/model.py"
)
assert "ls part/*.stl" in first_status_payload["validation_command"]
assert (
"(ls part/*.step || ls part/*.stp)"
in first_status_payload["validation_command"]
)
assert "STATUS: SUCCESS" in result
assert "Latest code" in result
response_payload = emitter.events[4][1]
assert response_payload["delta_text"] == "STATUS: SUCCESS\nSUMMARY: ok"
@pytest.mark.asyncio
async def test_cad_code_generator_retries_when_first_attempt_does_not_produce_code(
monkeypatch,
):
monkeypatch.chdir(WORKSPACE_ROOT)
calls = {"count": 0}
async def fake_specialist(**kwargs):
calls["count"] += 1
if calls["count"] == 1:
assert kwargs["history"] == []
assert "Current working directory:" in kwargs["message"]
assert (
"validation_command: uv run python part/model.py" in kwargs["message"]
)
assert "references/docs/api/README.md" in kwargs["message"]
yield ResponseYield(
response=(
"STATUS: SUCCESS\n"
"SUMMARY: planned but no file write\n"
"```python\n"
"print('draft only')\n"
"```"
),
messages=[],
)
return
assert "must write the required Python script directly" in kwargs["message"]
assert "part/model.py" in kwargs["message"]
assert kwargs["history"][0]["role"] == "user"
assert "target_file: part/model.py" in kwargs["history"][0]["content"]
assert kwargs["history"][1]["role"] == "assistant"
assert "planned but no file write" in kwargs["history"][1]["content"]
yield ResponseYield(
response="STATUS: SUCCESS\nCODE_WRITTEN: YES\nSUMMARY: wrote code",
messages=[],
)
monkeypatch.setattr(
code_tools_module,
"cad_code_generator_specialist",
fake_specialist,
)
code_reads = {"count": 0}
def fake_read_latest_code(path):
code_reads["count"] += 1
if code_reads["count"] == 1:
return None
return "print('ok')\n"
monkeypatch.setattr(code_tools_module, "_read_latest_code", fake_read_latest_code)
result = await cad_code_generator(
task="Create a cube as a new file. This is a create-new-file task.",
target_file_path="part/model.py",
event_emitter=None,
)
assert calls["count"] == 2
assert "CODE_WRITTEN: YES" in result
assert "Latest code" in result
@@ -0,0 +1,74 @@
import os
import subprocess
import sys
from types import SimpleNamespace
import pytest
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
import tools.command_tools as command_tools_module
from tools.command_tools import EXECUTE_COMMAND_TIMEOUT_SECONDS, execute_command
@pytest.mark.asyncio
async def test_execute_command_uses_600_second_timeout(monkeypatch):
captured = {}
async def fake_to_thread(func, *args, **kwargs):
captured["func"] = func
captured["args"] = args
captured["kwargs"] = kwargs
return SimpleNamespace(returncode=0, stdout="ok\n", stderr="")
monkeypatch.setattr(command_tools_module.asyncio, "to_thread", fake_to_thread)
result = await execute_command("uv run python demo.py")
assert result == "ok"
assert captured["args"][0] == "uv run python demo.py"
assert captured["kwargs"]["timeout"] == EXECUTE_COMMAND_TIMEOUT_SECONDS == 600
@pytest.mark.asyncio
async def test_execute_command_returns_english_failure_message(monkeypatch):
async def fake_to_thread(func, *args, **kwargs):
return SimpleNamespace(returncode=1, stdout="trace line\n", stderr="boom\n")
monkeypatch.setattr(command_tools_module.asyncio, "to_thread", fake_to_thread)
result = await execute_command("uv run python broken.py")
assert "Command failed with exit code 1." in result
assert "STDOUT:\ntrace line" in result
assert "STDERR:\nboom" in result
assert "Timeout may be caused by the program waiting for input" not in result
@pytest.mark.asyncio
async def test_execute_command_returns_english_timeout_message(monkeypatch):
async def fake_to_thread(func, *args, **kwargs):
raise subprocess.TimeoutExpired(
cmd="uv run python slow.py",
timeout=EXECUTE_COMMAND_TIMEOUT_SECONDS,
output="still running\n",
stderr="waiting\n",
)
monkeypatch.setattr(command_tools_module.asyncio, "to_thread", fake_to_thread)
result = await execute_command("uv run python slow.py")
assert (
f"Command timed out after {EXECUTE_COMMAND_TIMEOUT_SECONDS} seconds." in result
)
assert (
"The process may be stuck, waiting for input, or simply taking too long."
in result
)
assert "Partial STDOUT:\nstill running" in result
assert "Partial STDERR:\nwaiting" in result
@@ -0,0 +1,450 @@
#!/usr/bin/env python3
"""
Context system basic feature validation script.
Validates core features directly without relying on pytest:
1. Redis storage
2. File persistence
3. Message management
4. Context manager
"""
import os
import sys
import json
import tempfile
import shutil
import asyncio
import redis
from typing import Dict, List, Optional, Any
from datetime import datetime, timedelta
# Add the project root directory to the Python path.
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, project_root)
from context.context_manager import ContextManager
from context.context import RedisFileContextBackend
from context.schemas import Message
from config.config import get_config
class ContextSystemTester:
"""Context system tester."""
def __init__(self):
"""Initialize the tester."""
self.test_context_dir = tempfile.mkdtemp(prefix="test_context_")
self.test_redis_db = 13 # Use a dedicated test database.
# Point the configuration to the test directory.
config = get_config()
config.CONTEXT_DIR = self.test_context_dir
# Create a Redis connection.
self.redis_client = redis.Redis(db=self.test_redis_db, decode_responses=True)
self.redis_client.flushdb()
self.test_results = []
def cleanup(self):
"""Clean up the test environment."""
self.redis_client.flushdb()
self.redis_client.close()
if os.path.exists(self.test_context_dir):
shutil.rmtree(self.test_context_dir)
def log_test(self, test_name: str, success: bool, message: str = ""):
"""Record a test result."""
status = "✓ PASS" if success else "✗ FAIL"
print(f"{status} {test_name}: {message}")
self.test_results.append({
"test": test_name,
"success": success,
"message": message
})
def test_redis_connection(self) -> bool:
"""Test the Redis connection."""
try:
assert self.redis_client.ping()
self.log_test("Redis connection", True, "Redis service is healthy")
return True
except Exception as e:
self.log_test("Redis connection", False, f"Redis connection failed: {e}")
return False
async def test_message_storage(self) -> bool:
"""Test message storage functionality."""
try:
context_id = "test_message_storage"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
)
# Create test messages.
messages = [
Message(role="system", content="You are a helpful assistant"),
Message(role="user", content="Hello"),
Message(role="assistant", content="Hello! How can I help you?")
]
# Store messages.
for message in messages:
await backend.store_message(message)
# Verify message count.
assert backend.get_message_count() == len(messages)
# Retrieve messages.
retrieved_messages = backend.retrieve_messages()
assert len(retrieved_messages) == len(messages)
# Verify message content.
for i, (original, retrieved) in enumerate(zip(messages, retrieved_messages)):
assert original.role == retrieved.role
assert original.content == retrieved.content
self.log_test("Message storage", True, f"Successfully stored and retrieved {len(messages)} messages")
return True
except Exception as e:
self.log_test("Message storage", False, f"Message storage test failed: {e}")
return False
async def test_file_persistence(self) -> bool:
"""Test file persistence functionality."""
try:
context_id = "test_file_persistence"
file_path = os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=file_path
)
# Add messages.
messages = [
Message(role="user", content="Persistence test message 1"),
Message(role="assistant", content="Persistence test reply 1"),
Message(role="user", content="Persistence test message 2")
]
for message in messages:
await backend.store_message(message)
# Persist to file.
success = await backend.persist()
assert success
assert os.path.exists(file_path)
# Verify file contents.
with open(file_path, 'r', encoding='utf-8') as f:
data = json.load(f)
assert data['context_id'] == context_id
assert len(data['messages']) == len(messages)
self.log_test("File persistence", True, "Successfully persisted to the file system")
return True
except Exception as e:
self.log_test("File persistence", False, f"File persistence test failed: {e}")
return False
async def test_file_restoration(self) -> bool:
"""Test file restoration functionality."""
try:
context_id = "test_file_restoration"
file_path = os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
# Create the first backend and add data.
backend1 = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=file_path
)
messages = [
Message(role="user", content="Restoration test message 1"),
Message(role="assistant", content="Restoration test reply 1"),
Message(role="user", content="Restoration test message 2")
]
for message in messages:
await backend1.store_message(message)
# Persist data.
await backend1.persist()
# Clear Redis data to simulate a restart.
self.redis_client.flushdb()
# Create a new backend instance to simulate restart.
backend2 = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=file_path
)
# Restore from file.
success = await backend2.restore()
assert success
# Verify restored data.
restored_messages = backend2.retrieve_messages()
assert len(restored_messages) == len(messages)
for i, (original, restored) in enumerate(zip(messages, restored_messages)):
assert original.role == restored.role
assert original.content == restored.content
self.log_test("File restoration", True, "Successfully restored data from file")
return True
except Exception as e:
self.log_test("File restoration", False, f"File restoration test failed: {e}")
return False
async def test_context_manager(self) -> bool:
"""Test the context manager."""
try:
manager = ContextManager(backend_class=RedisFileContextBackend)
# Create a context.
context_id = "test_manager_context"
context = manager.create_context(context_id=context_id)
assert context is not None
assert context.context_id == context_id
# Add a message through the convenience interface.
success = await manager.add_message(
context_id=context_id,
role="user",
content="Message added through the manager"
)
assert success
# Get history.
history = manager.get_history(context_id)
assert len(history) == 1
assert history[0].role == "user"
assert history[0].content == "Message added through the manager"
# Test retrieving the context.
retrieved_context = manager.get_context(context_id)
assert retrieved_context is not None
assert retrieved_context.context_id == context_id
self.log_test("Context manager", True, "Context manager works correctly")
return True
except Exception as e:
self.log_test("Context manager", False, f"Context manager test failed: {e}")
return False
async def test_redis_persistence_verification(self) -> bool:
"""Verify data persistence in Redis."""
try:
context_id = "test_redis_verification"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
)
# Add messages.
messages = [
Message(role="user", content="Redis verification message 1"),
Message(role="assistant", content="Redis verification reply 1"),
Message(role="user", content="Redis verification message 2")
]
for message in messages:
await backend.store_message(message)
# Verify data in Redis.
messages_key = f"context:{context_id}:messages"
metadata_key = f"context:{context_id}:metadata"
# Check message data.
redis_messages = self.redis_client.lrange(messages_key, 0, -1)
assert len(redis_messages) == 3
# Check metadata.
redis_metadata = self.redis_client.get(metadata_key)
assert redis_metadata is not None
# Verify message content.
for i, message_json in enumerate(redis_messages):
message_data = json.loads(message_json)
assert message_data['role'] == messages[i].role
assert message_data['content'] == messages[i].content
self.log_test("Redis persistence verification", True, "Redis data storage works correctly")
return True
except Exception as e:
self.log_test("Redis persistence verification", False, f"Redis persistence verification failed: {e}")
return False
async def test_concurrent_access(self) -> bool:
"""Test concurrent access."""
try:
import threading
import time
context_id = "test_concurrent"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
)
# Add messages concurrently.
def add_messages(thread_id: int, count: int):
for i in range(count):
message = Message(role="user", content=f"Thread {thread_id} message {i}")
# Use a new event loop.
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
loop.run_until_complete(backend.store_message(message))
finally:
loop.close()
time.sleep(0.01)
# Create multiple threads.
threads = []
for i in range(3):
thread = threading.Thread(target=add_messages, args=(i, 5))
threads.append(thread)
thread.start()
# Wait for all threads to finish.
for thread in threads:
thread.join()
# Verify that all messages were added.
messages = backend.retrieve_messages()
assert len(messages) == 15 # 3 threads * 5 messages.
self.log_test("Concurrent access", True, "Concurrent access is safe")
return True
except Exception as e:
self.log_test("Concurrent access", False, f"Concurrent access test failed: {e}")
return False
async def run_all_tests(self) -> bool:
"""Run all tests."""
print("=" * 60)
print("Starting Context system tests")
print("=" * 60)
tests = [
("Redis connection", self.test_redis_connection),
("Message storage", self.test_message_storage),
("File persistence", self.test_file_persistence),
("File restoration", self.test_file_restoration),
("Context manager", self.test_context_manager),
("Redis persistence verification", self.test_redis_persistence_verification),
("Concurrent access", self.test_concurrent_access),
]
all_passed = True
for test_name, test_func in tests:
print(f"\nRunning test: {test_name}")
print("-" * 40)
try:
if asyncio.iscoroutinefunction(test_func):
result = await test_func()
else:
result = test_func()
if not result:
all_passed = False
except Exception as e:
self.log_test(test_name, False, f"Test exception: {e}")
all_passed = False
# Output test summary.
print("\n" + "=" * 60)
print("Test Result Summary")
print("=" * 60)
passed_count = sum(1 for result in self.test_results if result["success"])
total_count = len(self.test_results)
for result in self.test_results:
status = "✓" if result["success"] else "✗"
print(f"{status} {result['test']}: {result['message']}")
print(f"\nTotal: {passed_count}/{total_count} tests passed")
if all_passed:
print("\n🎉 All tests passed!")
print("Context system feature verification succeeded:")
print(" ✓ Redis storage works correctly")
print(" ✓ File system persistence works correctly")
print(" ✓ Message management works correctly")
print(" ✓ Context manager works correctly")
print(" ✓ Concurrent access is safe")
else:
print("\n❌ Some tests failed")
return all_passed
async def main():
"""Main function."""
print("Context System Basic Feature Validation")
print("=" * 40)
# Check Redis connection.
try:
r = redis.Redis(host="localhost", port=6379, decode_responses=True)
r.ping()
r.close()
print("✓ Redis service is available")
except Exception as e:
print(f"❌ Redis service is unavailable: {e}")
print("Please ensure the Redis service is running")
print("Startup commands:")
print(" macOS: brew services start redis")
print(" Linux: sudo systemctl start redis")
print(" Windows: redis-server")
return 1
# Run tests.
tester = ContextSystemTester()
try:
success = await tester.run_all_tests()
return 0 if success else 1
finally:
tester.cleanup()
if __name__ == "__main__":
exit(asyncio.run(main()))
@@ -0,0 +1,368 @@
#!/usr/bin/env python3
"""
Context system simplified test script.
Focuses on validating core features:
1. Redis storage
2. File persistence
3. Message management
"""
import os
import sys
import json
import tempfile
import shutil
import asyncio
import redis
import traceback
from typing import Dict, List, Optional, Any
# Add the project root directory to the Python path.
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, project_root)
from context.context import RedisFileContextBackend
from context.schemas import Message
class SimpleContextTester:
"""Simplified Context system tester."""
def __init__(self):
"""Initialize the tester."""
self.test_context_dir = tempfile.mkdtemp(prefix="test_context_")
self.test_redis_db = 12 # Use a dedicated test database.
# Create a Redis connection.
self.redis_client = redis.Redis(db=self.test_redis_db, decode_responses=True)
self.redis_client.flushdb()
self.test_results = []
def cleanup(self):
"""Clean up the test environment."""
self.redis_client.flushdb()
self.redis_client.close()
if os.path.exists(self.test_context_dir):
shutil.rmtree(self.test_context_dir)
def log_test(self, test_name: str, success: bool, message: str = ""):
"""Record a test result."""
status = "✓ PASS" if success else "✗ FAIL"
print(f"{status} {test_name}: {message}")
self.test_results.append({
"test": test_name,
"success": success,
"message": message
})
def test_redis_connection(self) -> bool:
"""Test the Redis connection."""
try:
assert self.redis_client.ping()
self.log_test("Redis connection", True, "Redis service is healthy")
return True
except Exception as e:
self.log_test("Redis connection", False, f"Redis connection failed: {e}")
return False
async def test_basic_message_storage(self) -> bool:
"""Test basic message storage functionality."""
try:
context_id = "test_basic_storage"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
)
# Create a test message.
message = Message(role="user", content="Test message")
# Store the message.
await backend.store_message(message)
# Verify message count.
count = backend.get_message_count()
assert count == 1, f"Expected 1 message, got {count}"
# Retrieve messages.
messages = backend.retrieve_messages()
assert len(messages) == 1, f"Expected 1 message, got {len(messages)}"
# Verify message content.
retrieved_message = messages[0]
assert retrieved_message.role == "user"
assert retrieved_message.content == "Test message"
self.log_test("Basic message storage", True, "Successfully stored and retrieved the message")
return True
except Exception as e:
error_msg = f"Basic message storage test failed: {e}\n{traceback.format_exc()}"
self.log_test("Basic message storage", False, error_msg)
return False
async def test_file_persistence_simple(self) -> bool:
"""Test simple file persistence functionality."""
try:
context_id = "test_file_simple"
file_path = os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=file_path
)
# Add one message.
message = Message(role="user", content="Persistence test message")
await backend.store_message(message)
# Persist to file.
success = await backend.persist()
assert success, "Persistence failed"
assert os.path.exists(file_path), f"File does not exist: {file_path}"
# Verify file contents.
with open(file_path, 'r', encoding='utf-8') as f:
data = json.load(f)
assert data['context_id'] == context_id
assert len(data['messages']) == 1
assert data['messages'][0]['content'] == "Persistence test message"
self.log_test("File persistence", True, "Successfully persisted to the file system")
return True
except Exception as e:
error_msg = f"File persistence test failed: {e}\n{traceback.format_exc()}"
self.log_test("File persistence", False, error_msg)
return False
async def test_redis_data_verification(self) -> bool:
"""Verify data in Redis."""
try:
context_id = "test_redis_verify"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
)
# Add a message.
message = Message(role="user", content="Redis verification message")
await backend.store_message(message)
# Ensure metadata is updated in Redis.
backend.update_metadata({"test_verification": "true"})
# Verify data in Redis.
messages_key = f"context:{context_id}:messages"
metadata_key = f"context:{context_id}:metadata"
# Check message data.
redis_messages = self.redis_client.lrange(messages_key, 0, -1)
assert len(redis_messages) == 1, f"Redis should contain 1 message, got {len(redis_messages)}"
# Check metadata.
redis_metadata = self.redis_client.get(metadata_key)
assert redis_metadata is not None, "Redis should contain metadata"
# Verify message content.
message_data = json.loads(redis_messages[0])
assert message_data['role'] == "user"
assert message_data['content'] == "Redis verification message"
# Verify metadata contents.
metadata_data = json.loads(redis_metadata)
assert metadata_data['context_id'] == context_id
assert metadata_data['test_verification'] == "true"
self.log_test("Redis data verification", True, "Redis data storage works correctly")
return True
except Exception as e:
error_msg = f"Redis data verification failed: {e}\n{traceback.format_exc()}"
self.log_test("Redis data verification", False, error_msg)
return False
async def test_metadata_management(self) -> bool:
"""Test metadata management."""
try:
context_id = "test_metadata"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
)
# Get initial metadata.
initial_metadata = backend.get_metadata()
assert initial_metadata['context_id'] == context_id
assert 'start_time' in initial_metadata
assert 'last_activity' in initial_metadata
# Update metadata.
new_metadata = {
'custom_field': 'custom_value',
'test_count': 42
}
backend.update_metadata(new_metadata)
# Verify the update.
updated_metadata = backend.get_metadata()
assert updated_metadata['custom_field'] == 'custom_value'
assert updated_metadata['test_count'] == 42
self.log_test("Metadata management", True, "Metadata management works correctly")
return True
except Exception as e:
error_msg = f"Metadata management test failed: {e}\n{traceback.format_exc()}"
self.log_test("Metadata management", False, error_msg)
return False
async def test_message_search(self) -> bool:
"""Test message search functionality."""
try:
context_id = "test_search"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host="localhost",
redis_port=6379,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
)
# Add messages containing specific keywords.
messages = [
Message(role="user", content="I want to learn Python programming"),
Message(role="assistant", content="Python is a great programming language"),
Message(role="user", content="Please tell me about machine learning")
]
for message in messages:
await backend.store_message(message)
# Search for messages containing "Python".
python_results = backend.search_messages("Python", limit=10)
assert len(python_results) == 2, f"Expected 2 messages containing Python, got {len(python_results)}"
# Search for messages containing "machine learning".
ml_results = backend.search_messages("machine learning", limit=10)
assert len(ml_results) == 1, f"Expected 1 message containing machine learning, got {len(ml_results)}"
self.log_test("Message search", True, "Message search works correctly")
return True
except Exception as e:
error_msg = f"Message search test failed: {e}\n{traceback.format_exc()}"
self.log_test("Message search", False, error_msg)
return False
async def run_all_tests(self) -> bool:
"""Run all tests."""
print("=" * 60)
print("Starting simplified Context system tests")
print("=" * 60)
tests = [
("Redis connection", self.test_redis_connection),
("Basic message storage", self.test_basic_message_storage),
("File persistence", self.test_file_persistence_simple),
("Redis data verification", self.test_redis_data_verification),
("Metadata management", self.test_metadata_management),
("Message search", self.test_message_search),
]
all_passed = True
for test_name, test_func in tests:
print(f"\nRunning test: {test_name}")
print("-" * 40)
try:
if asyncio.iscoroutinefunction(test_func):
result = await test_func()
else:
result = test_func()
if not result:
all_passed = False
except Exception as e:
error_msg = f"Test exception: {e}\n{traceback.format_exc()}"
self.log_test(test_name, False, error_msg)
all_passed = False
# Output test summary.
print("\n" + "=" * 60)
print("Test Result Summary")
print("=" * 60)
passed_count = sum(1 for result in self.test_results if result["success"])
total_count = len(self.test_results)
for result in self.test_results:
status = "✓" if result["success"] else "✗"
print(f"{status} {result['test']}: {result['message']}")
print(f"\nTotal: {passed_count}/{total_count} tests passed")
if all_passed:
print("\n🎉 All tests passed!")
print("Context system core feature verification succeeded:")
print(" ✓ Redis storage works correctly")
print(" ✓ File system persistence works correctly")
print(" ✓ Message management works correctly")
print(" ✓ Metadata management works correctly")
print(" ✓ Message search works correctly")
else:
print("\n❌ Some tests failed")
return all_passed
async def main():
"""Main function."""
print("Context System Simplified Feature Validation")
print("=" * 40)
# Check Redis connection.
try:
r = redis.Redis(host="localhost", port=6379, decode_responses=True)
r.ping()
r.close()
print("✓ Redis service is available")
except Exception as e:
print(f"❌ Redis service is unavailable: {e}")
print("Please ensure the Redis service is running")
print("Startup commands:")
print(" macOS: brew services start redis")
print(" Linux: sudo systemctl start redis")
print(" Windows: redis-server")
return 1
# Run tests.
tester = SimpleContextTester()
try:
success = await tester.run_all_tests()
return 0 if success else 1
finally:
tester.cleanup()
if __name__ == "__main__":
exit(asyncio.run(main()))
@@ -0,0 +1,764 @@
"""
Unit tests for the Context system.
Test coverage:
1. Redis storage
2. File system persistence
3. Message management
4. Context manager
5. Data serialization and deserialization
6. Automatic summary functionality
7. Metadata management
"""
import os
import json
import tempfile
import shutil
import asyncio
import pytest
import redis
from typing import Dict, List, Optional, Any
from unittest.mock import Mock, AsyncMock, patch
from datetime import datetime, timedelta
# Import the modules under test.
import sys
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from context.context_manager import ContextManager
from context.context import RedisFileContextBackend, ContextBackend
from context.schemas import Message, SketchPadItem
from config.config import get_config
def _make_configured_backend(redis_host: str, redis_port: int, redis_db: int):
class ConfiguredRedisFileContextBackend(RedisFileContextBackend):
def __init__(
self,
context_id: str,
llm_interface=None,
max_history_length: int = 5,
auto_summarize_trigger: int = 1000000,
file_path: Optional[str] = None,
redis_host: str = redis_host,
redis_port: int = redis_port,
redis_db: int = redis_db,
):
super().__init__(
context_id=context_id,
llm_interface=llm_interface,
max_history_length=max_history_length,
auto_summarize_trigger=auto_summarize_trigger,
redis_host=redis_host,
redis_port=redis_port,
redis_db=redis_db,
file_path=file_path,
)
return ConfiguredRedisFileContextBackend
class TestContextSystem:
"""Context system integration tests."""
@pytest.fixture(autouse=True)
def setup_and_teardown(self):
"""Set up and clean up around each test."""
# Set up the test environment.
self.test_context_dir = tempfile.mkdtemp(prefix="test_context_")
self.test_redis_db = 15 # Use a dedicated test database.
# Create test configuration.
self.original_config = None
if hasattr(get_config(), "CONTEXT_DIR"):
self.original_config = get_config().CONTEXT_DIR
# Point the configuration to the test directory.
config = get_config()
config.CONTEXT_DIR = self.test_context_dir
self.redis_host = config.REDIS_HOST
self.redis_port = int(config.REDIS_PORT)
self.backend_class = _make_configured_backend(
self.redis_host,
self.redis_port,
self.test_redis_db,
)
ContextManager._instance = None
# Create a Redis connection.
self.redis_client = redis.Redis(
host=self.redis_host,
port=self.redis_port,
db=self.test_redis_db,
decode_responses=True,
)
# Clear test data.
self.redis_client.flushdb()
yield
# Clean up the test environment.
self.redis_client.flushdb()
self.redis_client.close()
if os.path.exists(self.test_context_dir):
shutil.rmtree(self.test_context_dir)
ContextManager._instance = None
# Restore the original configuration.
if self.original_config:
config.CONTEXT_DIR = self.original_config
def test_redis_connection(self):
"""Test the Redis connection."""
assert self.redis_client.ping()
print("✓ Redis connection is healthy")
def test_context_backend_creation(self):
"""Test context backend creation."""
context_id = "test_context_001"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json"),
)
assert backend.context_id == context_id
assert backend.redis_client is not None
assert backend.file_path is not None
print("✓ Context backend created successfully")
@pytest.mark.asyncio
async def test_message_storage_and_retrieval(self):
"""Test message storage and retrieval."""
context_id = "test_context_002"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json"),
)
# Create test messages.
messages = [
Message(role="system", content="You are a helpful assistant"),
Message(role="user", content="Hello"),
Message(role="assistant", content="Hello! How can I help you?"),
Message(role="user", content="Please introduce Python"),
Message(role="assistant", content="Python is a high-level programming language..."),
]
# Store messages.
for message in messages:
await backend.store_message(message)
# Verify message count.
assert backend.get_message_count() == len(messages)
print(f"✓ Successfully stored {len(messages)} messages")
# Retrieve messages.
retrieved_messages = backend.retrieve_messages()
assert len(retrieved_messages) == len(messages)
# Verify message content.
for i, (original, retrieved) in enumerate(zip(messages, retrieved_messages)):
assert original.role == retrieved.role
assert original.content == retrieved.content
assert retrieved.timestamp is not None
print("✓ Message retrieval works correctly")
@pytest.mark.asyncio
async def test_file_persistence(self):
"""Test file persistence functionality."""
context_id = "test_context_003"
file_path = os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=file_path,
)
# Add some messages.
messages = [
Message(role="user", content="Test message 1"),
Message(role="assistant", content="Test reply 1"),
Message(role="user", content="Test message 2"),
Message(role="assistant", content="Test reply 2"),
]
for message in messages:
await backend.store_message(message)
# Persist to file.
success = await backend.persist()
assert success
assert os.path.exists(file_path)
# Verify file contents.
with open(file_path, "r", encoding="utf-8") as f:
data = json.load(f)
assert data["context_id"] == context_id
assert "messages" in data
assert "metadata" in data
assert len(data["messages"]) == len(messages)
print("✓ File persistence works correctly")
@pytest.mark.asyncio
async def test_file_restoration(self):
"""Test file restoration functionality."""
context_id = "test_context_004"
file_path = os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
# Create the first backend and add data.
backend1 = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=file_path,
)
messages = [
Message(role="user", content="Persistence test message 1"),
Message(role="assistant", content="Persistence test reply 1"),
Message(role="user", content="Persistence test message 2"),
]
for message in messages:
await backend1.store_message(message)
# Persist data.
await backend1.persist()
# Create a new backend instance to simulate restart.
backend2 = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=file_path,
)
# Restore from file.
success = await backend2.restore()
assert success
# Verify restored data.
restored_messages = backend2.retrieve_messages()
assert len(restored_messages) == len(messages)
for i, (original, restored) in enumerate(zip(messages, restored_messages)):
assert original.role == restored.role
assert original.content == restored.content
print("✓ File restoration works correctly")
@pytest.mark.asyncio
async def test_context_manager_integration(self):
"""Test context manager integration functionality."""
# Create a context manager.
manager = ContextManager(backend_class=self.backend_class)
# Create a context.
context_id = "test_context_005"
context = manager.create_context(context_id=context_id, max_history_length=10)
assert context is not None
assert context.context_id == context_id
# Add a message through the convenience interface.
success = await manager.add_message(
context_id=context_id,
message=Message(role="user", content="Message added through the manager"),
)
assert success
# Get history.
history = manager.get_history(context_id)
assert len(history) == 1
assert history[0].role == "user"
assert history[0].content == "Message added through the manager"
# Test retrieving the context.
retrieved_context = manager.get_context(context_id)
assert retrieved_context is not None
assert retrieved_context.context_id == context_id
print("✓ Context manager integration works correctly")
def test_metadata_management(self):
"""Test metadata management."""
context_id = "test_context_006"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json"),
)
# Get initial metadata.
initial_metadata = backend.get_metadata()
assert initial_metadata["context_id"] == context_id
assert "start_time" in initial_metadata
assert "last_activity" in initial_metadata
# Update metadata.
new_metadata = {"custom_field": "custom_value", "test_count": 42}
backend.update_metadata(new_metadata)
# Verify the update.
updated_metadata = backend.get_metadata()
assert updated_metadata["custom_field"] == "custom_value"
assert updated_metadata["test_count"] == 42
print("✓ Metadata management works correctly")
@pytest.mark.asyncio
async def test_message_search(self):
"""Test message search functionality."""
context_id = "test_context_007"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json"),
)
# Add messages containing specific keywords.
messages = [
Message(role="user", content="I want to learn Python programming"),
Message(role="assistant", content="Python is a great programming language"),
Message(role="user", content="Please tell me about machine learning"),
Message(role="assistant", content="Machine learning is a branch of artificial intelligence"),
Message(role="user", content="Python is commonly used in machine learning"),
]
for message in messages:
await backend.store_message(message)
# Search for messages containing "Python".
python_results = backend.search_messages("Python", limit=10)
assert len(python_results) == 3
# Search for messages containing "machine learning".
ml_results = backend.search_messages("machine learning", limit=10)
assert len(ml_results) == 3
# Search for a nonexistent keyword.
empty_results = backend.search_messages("nonexistent keyword", limit=10)
assert len(empty_results) == 0
print("✓ Message search works correctly")
@pytest.mark.asyncio
async def test_message_limit_management(self):
"""Test message count limit management."""
context_id = "test_context_008"
max_history_length = 3
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json"),
max_history_length=max_history_length,
)
# Add more messages than the limit allows.
for i in range(5):
message = Message(role="user", content=f"Message {i}")
await backend.store_message(message)
# Verify that only the latest messages are retained.
messages = backend.retrieve_messages()
assert len(messages) == max_history_length
# Verify that the retained messages are the latest ones.
expected_contents = ["Message 2", "Message 3", "Message 4"]
for i, message in enumerate(messages):
assert message.content == expected_contents[i]
print("✓ Message count limit management works correctly")
@pytest.mark.asyncio
async def test_context_serialization(self):
"""Test context serialization functionality."""
context_id = "test_context_009"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json"),
)
# Add messages and metadata.
messages = [
Message(role="user", content="Serialization test message"),
Message(role="assistant", content="Serialization test reply"),
]
for message in messages:
await backend.store_message(message)
backend.update_summary("This is a test conversation")
backend.update_metadata({"test_key": "test_value"})
# Serialize.
serialized_data = backend.serialize()
# Verify serialized data.
assert serialized_data["context_id"] == context_id
assert len(serialized_data["messages"]) == 2
assert serialized_data["summary"] == "This is a test conversation"
assert "serialization_timestamp" in serialized_data
# Create a new backend and deserialize.
new_backend = RedisFileContextBackend(
context_id="new_context",
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, "ctx_new_context.json"),
)
new_backend.deserialize(serialized_data)
# Verify deserialization result.
restored_messages = new_backend.retrieve_messages()
assert len(restored_messages) == 2
assert new_backend.get_summary() == "This is a test conversation"
print("✓ Context serialization works correctly")
@pytest.mark.asyncio
async def test_context_manager_list_and_delete(self):
"""Test context manager list and delete functionality."""
manager = ContextManager(backend_class=self.backend_class)
# Create multiple contexts.
context_ids = ["test_ctx_001", "test_ctx_002", "test_ctx_003"]
for context_id in context_ids:
context = manager.create_context(context_id=context_id)
await manager.add_message(
context_id,
Message(role="user", content=f"Test message {context_id}"),
)
# List all contexts.
contexts = manager.list_contexts()
assert len(contexts) >= len(context_ids)
# Verify that the created contexts are all listed.
found_contexts = [ctx["context_id"] for ctx in contexts]
for context_id in context_ids:
assert context_id in found_contexts
# Delete one context.
delete_success = manager.delete_context(context_ids[0])
assert delete_success
# Verify it cannot be retrieved after deletion.
deleted_context = manager.get_context(context_ids[0])
assert deleted_context is None
print("✓ Context manager list and delete functionality works correctly")
@pytest.mark.asyncio
async def test_redis_persistence_verification(self):
"""Verify data persistence in Redis."""
context_id = "test_redis_persistence"
# Create a backend and add data.
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json"),
)
# Add messages.
messages = [
Message(role="user", content="Redis persistence test message 1"),
Message(role="assistant", content="Redis persistence test reply 1"),
Message(role="user", content="Redis persistence test message 2"),
]
for message in messages:
await backend.store_message(message)
# Verify data in Redis.
messages_key = f"context:{context_id}:messages"
metadata_key = f"context:{context_id}:metadata"
summary_key = f"context:{context_id}:summary"
# Check message data.
redis_messages = self.redis_client.lrange(messages_key, 0, -1)
assert len(redis_messages) == 3
# Check metadata.
redis_metadata = self.redis_client.get(metadata_key)
assert redis_metadata is not None
# Verify message content.
for i, message_json in enumerate(redis_messages):
message_data = json.loads(message_json)
expected_message = messages[-(i + 1)]
assert message_data["role"] == expected_message.role
assert message_data["content"] == expected_message.content
print("✓ Redis data persistence verification passed")
@pytest.mark.asyncio
async def test_file_system_persistence_verification(self):
"""Verify data persistence in the file system."""
context_id = "test_file_persistence"
file_path = os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=file_path,
)
# Add messages and metadata.
messages = [
Message(role="user", content="File persistence test message 1"),
Message(role="assistant", content="File persistence test reply 1"),
Message(role="user", content="File persistence test message 2"),
]
for message in messages:
await backend.store_message(message)
backend.update_summary("File persistence test summary")
backend.update_metadata({"file_test": "file_value"})
# Persist to file.
await backend.persist()
# Verify file exists.
assert os.path.exists(file_path)
# Read and verify file contents.
with open(file_path, "r", encoding="utf-8") as f:
file_data = json.load(f)
# Verify file structure.
assert file_data["context_id"] == context_id
assert len(file_data["messages"]) == 3
assert file_data["summary"] == "File persistence test summary"
assert file_data["metadata"]["file_test"] == "file_value"
# Verify message content.
for i, message_data in enumerate(file_data["messages"]):
assert message_data["role"] == messages[i].role
assert message_data["content"] == messages[i].content
print("✓ File system data persistence verification passed")
@pytest.mark.asyncio
async def test_concurrent_access(self):
"""Test concurrent access."""
import threading
import time
context_id = "test_concurrent"
backend = RedisFileContextBackend(
context_id=context_id,
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path=os.path.join(self.test_context_dir, f"ctx_{context_id}.json"),
max_history_length=20,
)
# Add messages concurrently.
def add_messages(thread_id: int, count: int):
for i in range(count):
message = Message(role="user", content=f"Thread {thread_id} message {i}")
asyncio.run(backend.store_message(message))
time.sleep(0.01) # Small delay.
# Create multiple threads.
threads = []
for i in range(3):
thread = threading.Thread(target=add_messages, args=(i, 5))
threads.append(thread)
thread.start()
# Wait for all threads to finish.
for thread in threads:
thread.join()
# Verify that all messages were added.
messages = backend.retrieve_messages()
assert len(messages) == 15 # 3 threads * 5 messages.
print("✓ Concurrent access test passed")
def test_error_handling(self):
"""Test error handling."""
# Test an invalid Redis connection.
try:
invalid_backend = RedisFileContextBackend(
context_id="test_error",
redis_host="invalid_host",
redis_port=9999,
redis_db=0,
file_path=os.path.join(self.test_context_dir, "ctx_test_error.json"),
)
# If the connection fails, it should raise an exception.
invalid_backend.redis_client.ping()
except Exception as e:
print(f"✓ Expected Redis connection error: {type(e).__name__}")
# Test an invalid file path.
backend = RedisFileContextBackend(
context_id="test_error_file",
redis_host=self.redis_host,
redis_port=self.redis_port,
redis_db=self.test_redis_db,
file_path="/invalid/path/ctx_test.json",
)
# Try to persist to an invalid path.
async def test_invalid_persist():
return await backend.persist()
result = asyncio.run(test_invalid_persist())
assert not result # Should fail.
print("✓ Error handling test passed")
class TestContextManagerAdvanced:
"""Advanced ContextManager feature tests."""
@pytest.fixture(autouse=True)
def setup(self):
"""Set up the test environment."""
self.test_context_dir = tempfile.mkdtemp(prefix="test_manager_")
self.test_redis_db = 14
# Modify configuration.
config = get_config()
config.CONTEXT_DIR = self.test_context_dir
self.redis_host = config.REDIS_HOST
self.redis_port = int(config.REDIS_PORT)
self.backend_class = _make_configured_backend(
self.redis_host,
self.redis_port,
self.test_redis_db,
)
ContextManager._instance = None
# Create a Redis connection.
self.redis_client = redis.Redis(
host=self.redis_host,
port=self.redis_port,
db=self.test_redis_db,
decode_responses=True,
)
self.redis_client.flushdb()
yield
# Clean up.
self.redis_client.flushdb()
self.redis_client.close()
shutil.rmtree(self.test_context_dir)
ContextManager._instance = None
@pytest.mark.asyncio
async def test_context_manager_singleton(self):
"""Test the ContextManager singleton pattern."""
manager1 = ContextManager(backend_class=self.backend_class)
manager2 = ContextManager(backend_class=self.backend_class)
assert manager1 is manager2
print("✓ ContextManager singleton pattern works correctly")
@pytest.mark.asyncio
async def test_context_manager_bulk_operations(self):
"""Test ContextManager batch operations."""
manager = ContextManager(backend_class=self.backend_class)
# Create multiple contexts.
context_ids = ["bulk_test_001", "bulk_test_002", "bulk_test_003"]
contexts = []
for context_id in context_ids:
context = manager.create_context(context_id=context_id)
contexts.append(context)
# Add a message.
await manager.add_message(
context_id,
Message(role="user", content=f"Batch test message {context_id}"),
)
# Save all contexts.
saved_count = await manager.save_all_contexts()
assert saved_count == len(context_ids)
# Verify files exist.
for context_id in context_ids:
file_path = os.path.join(self.test_context_dir, f"ctx_{context_id}.json")
assert os.path.exists(file_path)
print("✓ ContextManager batch operations work correctly")
@pytest.mark.asyncio
async def test_context_manager_cleanup(self):
"""Test ContextManager cleanup functionality."""
manager = ContextManager(backend_class=self.backend_class)
# Create a context and add a message.
context_id = "cleanup_test"
context = manager.create_context(context_id=context_id)
await manager.add_message(
context_id,
Message(role="user", content="Cleanup test message"),
)
# Simulate long-term inactivity by modifying last_activity in metadata.
context.update_metadata(
{"last_activity": (datetime.now() - timedelta(hours=2)).isoformat()}
)
# Run cleanup with a short inactivity threshold.
cleaned_count = await manager.cleanup_inactive_contexts(max_inactive_time=1)
assert cleaned_count == 1
# Verify the context was removed from the active cache but can still be restored from persistent storage on demand.
assert context_id not in manager._active_contexts
retrieved_context = manager.get_context(context_id)
assert retrieved_context is not None
print("✓ ContextManager cleanup functionality works correctly")
if __name__ == "__main__":
# Run tests.
pytest.main([__file__, "-v", "-s"])
@@ -0,0 +1,144 @@
import fnmatch
import os
import sys
import threading
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
class _FakeRedis:
_store: dict[str, object] = {}
def __init__(self, *args, **kwargs):
pass
@classmethod
def reset(cls) -> None:
cls._store = {}
def keys(self, pattern: str):
return [key for key in self._store.keys() if fnmatch.fnmatch(key, pattern)]
def delete(self, *keys: str):
deleted = 0
for key in keys:
if key in self._store:
deleted += 1
del self._store[key]
return deleted
def test_delete_context_removes_file_and_redis_keys(tmp_path, monkeypatch):
import context.context_manager as context_manager_module
_FakeRedis.reset()
monkeypatch.setattr(context_manager_module.redis, "Redis", _FakeRedis)
context_id = "conv-1"
context_file = tmp_path / f"ctx_{context_id}.json"
context_file.write_text("{}", encoding="utf-8")
_FakeRedis._store = {
f"context:{context_id}:messages": ["message"],
f"context:{context_id}:metadata": "{}",
}
manager = object.__new__(context_manager_module.ContextManager)
manager._lock = threading.RLock()
manager._active_contexts = {}
manager.context_dir = str(tmp_path)
manager.config = type(
"Cfg",
(),
{"REDIS_HOST": "localhost", "REDIS_PORT": 9736, "REDIS_DB": 0},
)()
assert manager.delete_context(context_id) is True
assert context_file.exists() is False
assert not any(
key.startswith(f"context:{context_id}:") for key in _FakeRedis._store
)
def test_delete_sketch_pad_removes_file_and_redis_keys(tmp_path, monkeypatch):
import context.sketch_manager as sketch_manager_module
_FakeRedis.reset()
monkeypatch.setattr(sketch_manager_module.redis, "Redis", _FakeRedis)
sketch_id = "conv-1"
sketch_file = tmp_path / f"skt_{sketch_id}.json"
sketch_file.write_text("{}", encoding="utf-8")
_FakeRedis._store = {
f"sketch_pad:{sketch_id}:code": "print('x')",
f"sketch_pad:{sketch_id}:tag:model": ["code"],
}
manager = object.__new__(sketch_manager_module.SketchManager)
manager._lock = threading.RLock()
manager._active_sketches = {}
manager.sketch_dir = str(tmp_path)
manager.config = type(
"Cfg",
(),
{"REDIS_HOST": "localhost", "REDIS_PORT": 9736, "REDIS_DB": 0},
)()
assert manager.delete_sketch_pad(sketch_id) is True
assert sketch_file.exists() is False
assert not any(
key.startswith(f"sketch_pad:{sketch_id}:") for key in _FakeRedis._store
)
def test_delete_all_conversations_discovers_ids_across_sources(tmp_path):
from context.conversation_manager import ConversationManager
marker_id = "marker-only"
active_id = "active-only"
context_id = "context-only"
sketch_id = "sketch-only"
conversations_dir = tmp_path / "conversations"
conversations_dir.mkdir()
(conversations_dir / f"conv_{marker_id}.marker").write_text("", encoding="utf-8")
deleted_contexts: list[str] = []
deleted_sketches: list[str] = []
class _FakeContextManager:
def list_context_ids(self):
return [context_id]
def delete_context(self, conversation_id: str):
deleted_contexts.append(conversation_id)
return conversation_id in {marker_id, active_id, context_id, sketch_id}
class _FakeSketchManager:
def list_sketch_ids(self):
return [sketch_id]
def delete_sketch_pad(self, conversation_id: str):
deleted_sketches.append(conversation_id)
return conversation_id in {marker_id, active_id, context_id, sketch_id}
manager = object.__new__(ConversationManager)
manager._lock = threading.RLock()
manager._active_conversations = {active_id: object()}
manager.context_manager = _FakeContextManager()
manager.sketch_manager = _FakeSketchManager()
manager.conversations_dir = str(conversations_dir)
deleted_ids = manager.delete_all_conversations()
assert sorted(deleted_ids) == sorted([marker_id, active_id, context_id, sketch_id])
assert sorted(deleted_contexts) == sorted(
[marker_id, active_id, context_id, sketch_id]
)
assert sorted(deleted_sketches) == sorted(
[marker_id, active_id, context_id, sketch_id]
)
@@ -0,0 +1,251 @@
# pyright: reportAssignmentType=false, reportArgumentType=false, reportIndexIssue=false
import os
import sys
from typing import Any, cast
import pytest
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
from context.context import RedisFileContextBackend
from context.schemas import Message
from web_interface.models import (
ChatMessageContentImage,
ChatMessageContentImageUrl,
ChatMessageContentText,
)
from web_interface.routers.conversation_router import get_conversation_history
class _FakeRedis:
_store: dict[str, object] = {}
def __init__(self, *args, **kwargs):
pass
@classmethod
def reset(cls) -> None:
cls._store = {}
def lpush(self, key: str, value: str) -> None:
values = cast(list[str], self._store.setdefault(key, []))
values.insert(0, value)
def rpush(self, key: str, value: str) -> None:
values = cast(list[str], self._store.setdefault(key, []))
values.append(value)
def lrange(self, key: str, start: int, end: int):
values = list(cast(list[str], self._store.get(key, [])))
if end == -1:
end = len(values) - 1
return values[start : end + 1]
def llen(self, key: str) -> int:
values = self._store.get(key, [])
assert isinstance(values, list)
return len(values)
def delete(self, key: str) -> None:
self._store.pop(key, None)
def set(self, key: str, value: str) -> None:
self._store[key] = value
def get(self, key: str):
return self._store.get(key)
def ltrim(self, key: str, start: int, end: int) -> None:
values = list(cast(list[str], self._store.get(key, [])))
if end == -1:
trimmed = values[start:]
else:
trimmed = values[start : end + 1]
self._store[key] = trimmed
@pytest.mark.asyncio
async def test_full_history_persists_even_when_working_memory_summarizes(
tmp_path, monkeypatch
):
import context.context as context_module
_FakeRedis.reset()
monkeypatch.setattr(context_module.redis, "Redis", _FakeRedis)
context_file = tmp_path / "ctx_conv-1.json"
backend = RedisFileContextBackend(
context_id="conv-1",
llm_interface=None,
max_history_length=2,
auto_summarize_trigger=2,
file_path=str(context_file),
)
cast(Any, backend).llm_interface = object()
async def fake_summarize(messages):
return "summary"
backend._summarize_func = fake_summarize
await backend.store_message(Message(role="user", content="first"))
await backend.store_message(Message(role="assistant", content="second"))
await backend.store_message(Message(role="user", content="third"))
assert [message.content for message in backend.retrieve_messages()] == ["third"]
assert [message.content for message in backend.retrieve_full_messages()] == [
"first",
"second",
"third",
]
await backend.persist()
_FakeRedis.reset()
restored = RedisFileContextBackend(
context_id="conv-1",
llm_interface=None,
max_history_length=2,
auto_summarize_trigger=2,
file_path=str(context_file),
)
await restored.restore()
assert [message.content for message in restored.retrieve_messages()] == ["third"]
assert [message.content for message in restored.retrieve_full_messages()] == [
"first",
"second",
"third",
]
@pytest.mark.asyncio
async def test_conversation_history_endpoint_returns_full_persisted_history():
archived_messages = [
Message(role="user", content="first"),
Message(role="assistant", content="second"),
Message(role="user", content="third"),
]
working_messages = [archived_messages[-1]]
class _FakeContext:
def retrieve_messages(self):
return working_messages
def retrieve_full_messages(self):
return archived_messages
class _FakeConversation:
uuid = "conversation-1"
context = _FakeContext()
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
class _FakeConversationManager:
def get_conversation(self, conversation_id):
return _FakeConversation() if conversation_id == "conversation-1" else None
class _FakeState:
conversation_manager = _FakeConversationManager()
response = cast(
Any,
await get_conversation_history(
conversation_id="conversation-1",
limit=None,
state=cast(Any, _FakeState()),
),
)
assert response["total_messages"] == 3
assert [message["content"] for message in response["messages"]] == [
"first",
"second",
"third",
]
@pytest.mark.asyncio
async def test_large_auto_summarize_trigger_effectively_disables_summary(
tmp_path, monkeypatch
):
import context.context as context_module
_FakeRedis.reset()
monkeypatch.setattr(context_module.redis, "Redis", _FakeRedis)
context_file = tmp_path / "ctx_conv-2.json"
backend = RedisFileContextBackend(
context_id="conv-2",
llm_interface=None,
max_history_length=2,
auto_summarize_trigger=999999,
file_path=str(context_file),
)
cast(Any, backend).llm_interface = object()
summarize_call_count = 0
async def fake_summarize(messages):
nonlocal summarize_call_count
summarize_call_count += 1
return "summary"
backend._summarize_func = fake_summarize
await backend.store_message(Message(role="user", content="first"))
await backend.store_message(Message(role="assistant", content="second"))
await backend.store_message(Message(role="user", content="third"))
assert summarize_call_count == 0
assert backend.get_summary() is None
assert [message.content for message in backend.retrieve_messages()] == [
"second",
"third",
]
def test_context_message_accepts_web_multimodal_models():
message = Message(
role="user",
content=[
ChatMessageContentText.model_validate(
{"type": "text", "text": "Treat this as an assembly"}
),
ChatMessageContentImage.model_validate(
{
"type": "image_url",
"image_url": ChatMessageContentImageUrl.model_validate(
{
"url": "data:image/png;base64,abcd",
"local_path": "/tmp/query_image_001.png",
}
),
}
),
],
)
assert isinstance(message.content, list)
payload = message.model_dump()
assert payload["content"] == [
{"type": "text", "text": "Treat this as an assembly"},
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,abcd",
"detail": None,
"local_path": "/tmp/query_image_001.png",
},
},
]
@@ -0,0 +1,43 @@
import os
import sys
from contextlib import contextmanager
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
from observability.langfuse_tracing import propagate_conversation_session
import observability.langfuse_tracing as tracing_module
def test_propagate_conversation_session_uses_conversation_id_as_session(monkeypatch):
captured: dict[str, object] = {}
@contextmanager
def fake_propagate_attributes(**kwargs):
captured.update(kwargs)
yield
monkeypatch.setattr(tracing_module, "_langfuse_is_configured", lambda: True)
monkeypatch.setattr(
tracing_module,
"propagate_attributes",
fake_propagate_attributes,
)
with propagate_conversation_session(
conversation_id="conversation-123",
metadata={"model": "cadagent", "turn": 2},
tags=["cadagent", "event_stream"],
):
pass
assert captured["session_id"] == "conversation-123"
assert captured["tags"] == ["cadagent", "event_stream"]
assert captured["metadata"] == {
"conversation_id": "conversation-123",
"model": "cadagent",
"turn": "2",
}
@@ -0,0 +1,339 @@
import os
import sys
from types import SimpleNamespace
import numpy as np
import pytest
from PIL import Image
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
import tools.model_view_tools as model_view_tools_module
import tools.reference_image as reference_image_module
from tools.model_view_tools import get_visual_feedback
@pytest.mark.asyncio
async def test_get_visual_feedback_uses_latest_uploaded_image_for_both_llm_steps(
monkeypatch, tmp_path
):
image_path = tmp_path / "query_image_001.png"
image_path.write_bytes(b"fake-png-bytes")
render_path = tmp_path / "render.png"
render_path.write_bytes(b"fake-render")
class FakeContext:
def retrieve_full_messages(self):
return [
SimpleNamespace(
role="user",
content=[
{"type": "text", "text": "Inspect this part"},
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,abcd",
"local_path": str(image_path),
},
},
],
)
]
monkeypatch.setattr(
reference_image_module,
"get_current_context",
lambda: FakeContext(),
)
monkeypatch.setattr(
model_view_tools_module, "_require_simplecad_renderer", lambda: None
)
monkeypatch.setattr(
model_view_tools_module, "print_tool_output", lambda *args, **kwargs: None
)
class FakeSketchPad:
async def set_item(self, key, value, ttl=None, summary=None, tags=None):
return key
monkeypatch.setattr(
model_view_tools_module,
"get_current_sketch_pad",
lambda: FakeSketchPad(),
)
captured = {}
async def fake_question_generator(user_query, code, query_image_path):
captured["question_image_path"] = (
str(query_image_path.path) if query_image_path else None
)
return "Checklist"
async def fake_visual_feedback_generator(
questions, multi_view_results, query_image_path
):
captured["visual_image_path"] = (
str(query_image_path.path) if query_image_path else None
)
captured["multi_view_results"] = str(multi_view_results.path)
return "Looks correct\nPASS"
monkeypatch.setattr(
model_view_tools_module,
"question_generator",
fake_question_generator,
)
monkeypatch.setattr(
model_view_tools_module,
"visual_feedback_generator",
fake_visual_feedback_generator,
)
from SimpleLLMFunc.type import ImgPath
monkeypatch.setattr(
model_view_tools_module,
"render_multi_view_model",
lambda model_path, output_path: ImgPath(render_path, detail="high"),
)
result = await get_visual_feedback(
user_query="Inspect this part",
code="result = None",
model_path="./part/model.stl",
)
expected_image_path = str(image_path.resolve())
assert captured["question_image_path"] == expected_image_path
assert captured["visual_image_path"] == expected_image_path
assert captured["multi_view_results"] == str(render_path.resolve())
assert "Model path: ./part/model.stl" in result
def test_camera_relative_light_rig_tracks_camera_direction() -> None:
for view_dir in (
np.array([1.0, 0.0, 0.0]),
np.array([-1.0, 0.0, 0.0]),
np.array([0.0, 0.0, 1.0]),
np.array([1.0, 1.0, -1.0]),
):
normalized_view = view_dir / np.linalg.norm(view_dir)
light_dirs, light_weights, ambient = (
model_view_tools_module._camera_relative_light_rig(normalized_view)
)
assert ambient > 0.0
assert len(light_dirs) == len(light_weights) >= 3
assert np.dot(light_dirs[0], normalized_view) > 0.85
assert all(
abs(np.linalg.norm(light_dir) - 1.0) < 1e-6 for light_dir in light_dirs
)
def test_camera_relative_shading_keeps_front_faces_bright_across_views() -> None:
base_color = np.array([0.72, 0.76, 0.81], dtype=float)
front_face_brightness: list[float] = []
for view_dir in (
np.array([1.0, 0.0, 0.0]),
np.array([-1.0, 0.0, 0.0]),
np.array([0.0, 1.0, 0.0]),
np.array([1.0, 1.0, -1.0]),
):
normalized_view = view_dir / np.linalg.norm(view_dir)
normals = np.stack([normalized_view, -normalized_view])
shaded = model_view_tools_module._shade_normals_camera_relative(
normals,
base_color,
normalized_view,
)
front_face_brightness.append(float(np.mean(shaded[0, :3])))
assert float(np.mean(shaded[0, :3])) > float(np.mean(shaded[1, :3]))
assert min(front_face_brightness) > 0.55
assert max(front_face_brightness) - min(front_face_brightness) < 0.12
def test_surface_shading_normals_keep_planar_triangles_consistent() -> None:
class FakeVector:
def __init__(self, x: float, y: float, z: float) -> None:
self.x = x
self.y = y
self.z = z
class FakeCadFace:
def geomType(self) -> str:
return "PLANE"
def normalAt(self, location=None):
return FakeVector(0.0, 0.0, 1.0)
class FakeFace:
cq_face = FakeCadFace()
tri_pts = np.array(
[
[[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [1.0, 1.0, 0.0]],
[[0.0, 0.0, 0.0], [1.0, 1.0, 0.0], [0.0, 1.0, 0.0]],
],
dtype=float,
)
normals = model_view_tools_module._compute_surface_shading_normals(
FakeFace(), tri_pts
)
assert normals.shape == (2, 3)
assert np.allclose(normals[0], [0.0, 0.0, 1.0])
assert np.allclose(normals[1], [0.0, 0.0, 1.0])
def test_direct_cad_renderer_is_used_for_brep_like_formats() -> None:
assert model_view_tools_module._should_use_direct_cad_renderer("part.step") is True
assert model_view_tools_module._should_use_direct_cad_renderer("part.stp") is True
assert model_view_tools_module._should_use_direct_cad_renderer("part.brep") is True
assert model_view_tools_module._should_use_direct_cad_renderer("part.bin") is True
assert model_view_tools_module._should_use_direct_cad_renderer("part.stl") is False
def test_prefer_cad_native_model_path_uses_step_over_stl(tmp_path) -> None:
stl_path = tmp_path / "model.stl"
step_path = tmp_path / "model.step"
stl_path.write_text("solid", encoding="utf-8")
step_path.write_text("step", encoding="utf-8")
selected = model_view_tools_module._prefer_cad_native_model_path(str(stl_path))
assert selected == str(step_path.resolve())
def test_prefer_cad_native_model_path_keeps_stl_without_step(tmp_path) -> None:
stl_path = tmp_path / "model.stl"
stl_path.write_text("solid", encoding="utf-8")
selected = model_view_tools_module._prefer_cad_native_model_path(str(stl_path))
assert selected == str(stl_path.resolve())
def test_load_renderable_shapes_flattens_step_compounds(monkeypatch) -> None:
class FakeCadShape:
def __init__(self, shape_type: str, solids=None) -> None:
self._shape_type = shape_type
self._solids = list(solids or [])
def ShapeType(self) -> str:
return self._shape_type
def Solids(self):
return list(self._solids)
class FakeWrappedSolid:
def __init__(self, obj) -> None:
self.obj = obj
class FakeWorkplane:
def __init__(self, values) -> None:
self._values = values
def vals(self):
return list(self._values)
solid_a = FakeCadShape("Solid")
solid_b = FakeCadShape("Solid")
compound = FakeCadShape("Compound", solids=[solid_a, solid_b])
fake_cq = SimpleNamespace(
importers=SimpleNamespace(
importShape=lambda import_type, model_path: FakeWorkplane([compound])
)
)
monkeypatch.setattr(model_view_tools_module, "cq", fake_cq)
monkeypatch.setattr(model_view_tools_module, "ScadSolid", FakeWrappedSolid)
monkeypatch.setattr(
model_view_tools_module, "_require_simplecad_renderer", lambda: None
)
result = model_view_tools_module._load_renderable_shapes("part.step")
assert [wrapped.obj for wrapped in result] == [solid_a, solid_b]
def test_feature_edge_mask_detects_normal_and_depth_discontinuities() -> None:
mask = np.ones((8, 8), dtype=bool)
depth = np.zeros((8, 8), dtype=float)
normals = np.zeros((8, 8, 3), dtype=float)
normals[:, :4] = np.array([0.0, 0.0, 1.0])
normals[:, 4:] = np.array([1.0, 0.0, 0.0])
edge_mask = model_view_tools_module._compute_feature_edge_mask(
mask,
depth,
normals,
depth_jump_threshold=10.0,
normal_cos_threshold=0.95,
)
assert edge_mask[:, 3:5].any()
depth[:, 4:] = 4.0
normals[:, :] = np.array([0.0, 0.0, 1.0])
edge_mask = model_view_tools_module._compute_feature_edge_mask(
mask,
depth,
normals,
depth_jump_threshold=1.0,
normal_cos_threshold=0.95,
)
assert edge_mask[:, 3:5].any()
def test_direct_rasterizer_supersamples_and_preserves_target_size() -> None:
triangles = [
np.array([[0.0, 0.0, 1.0], [1.0, 0.0, 1.0], [1.0, 1.0, 1.0]], dtype=float),
np.array([[0.0, 0.0, 1.0], [1.0, 1.0, 1.0], [0.0, 1.0, 1.0]], dtype=float),
]
normals = [
np.array([0.0, 0.0, 1.0], dtype=float),
np.array([0.0, 0.0, 1.0], dtype=float),
]
image = model_view_tools_module._rasterize_projected_triangles(
triangles,
normals,
image_size=(48, 48),
zoom=4.0,
background_rgb=np.array([255, 255, 255], dtype=np.uint8),
fill_rgb=np.array([180, 190, 200], dtype=np.uint8),
outline_rgb=np.array([0, 0, 0], dtype=np.uint8),
)
pixels = np.asarray(image)
unique_colors = np.unique(pixels.reshape(-1, 3), axis=0)
assert image.size == (48, 48)
assert len(unique_colors) > 3
def test_axis_triad_overlay_draws_small_corner_marker() -> None:
image = Image.new("RGB", (320, 320), "white")
annotated = model_view_tools_module._add_axis_triad_overlay(
image,
np.array([1.0, 1.0, 1.0]) / np.sqrt(3.0),
)
original = np.asarray(image)
updated = np.asarray(annotated)
diff = np.abs(updated.astype(int) - original.astype(int)).sum(axis=2)
changed_pixels = np.argwhere(diff > 0)
assert changed_pixels.size > 0
assert int(changed_pixels[:, 0].max()) > 220
assert int(changed_pixels[:, 1].max()) < 120
@@ -0,0 +1,155 @@
from __future__ import annotations
from pathlib import Path
from experiments.scripts.output_test_bench import (
CaseSpec,
build_progress_snapshot,
build_benchmark_prompt,
discover_cases,
load_existing_results,
merge_artifacts,
scan_output_directory,
)
def test_discover_cases_extracts_case_metadata(tmp_path: Path) -> None:
dataset_root = tmp_path / "output_test"
case_dir = dataset_root / "simple" / "0000" / "00000797"
case_dir.mkdir(parents=True)
(case_dir / "description.txt").write_text("Create a part", encoding="utf-8")
(case_dir / "00000797.png").write_bytes(b"png")
(case_dir / "00000797_cip.py").write_text("# cip", encoding="utf-8")
(case_dir / "00000797_cq.py").write_text("# cq", encoding="utf-8")
(case_dir / "00000797.step").write_text("step", encoding="utf-8")
cases = discover_cases(dataset_root)
assert len(cases) == 1
case = cases[0]
assert case.case_id == "simple/0000/00000797"
assert case.case_slug == "simple__0000__00000797"
assert Path(case.description_path).name == "description.txt"
assert Path(case.image_path or "").name == "00000797.png"
assert Path(case.reference_cip_path or "").name == "00000797_cip.py"
assert Path(case.reference_cq_path or "").name == "00000797_cq.py"
assert Path(case.reference_step_path or "").name == "00000797.step"
def test_discover_cases_works_when_dataset_root_is_already_simple_dir(
tmp_path: Path,
) -> None:
dataset_root = tmp_path / "simple"
case_dir = dataset_root / "0000" / "00000797"
case_dir.mkdir(parents=True)
(case_dir / "description.txt").write_text("Create a part", encoding="utf-8")
(case_dir / "00000797.png").write_bytes(b"png")
cases = discover_cases(dataset_root)
assert len(cases) == 1
case = cases[0]
assert case.case_id == "0000/00000797"
assert case.bucket_id == "0000"
assert case.sample_id == "00000797"
def test_build_benchmark_prompt_requires_full_automation(tmp_path: Path) -> None:
case = CaseSpec(
case_id="simple/0000/00000797",
case_slug="simple__0000__00000797",
bucket_id="0000",
sample_id="00000797",
case_dir=str(tmp_path),
description_path=str(tmp_path / "description.txt"),
image_path=str(tmp_path / "00000797.png"),
reference_cip_path=None,
reference_cq_path=None,
reference_step_path=None,
)
prompt = build_benchmark_prompt(
case=case,
description_text="Build a flange from the reference image.",
requested_output_dir=Path(
"/repo/workspace/experiments/output_test/r1/simple/0000/00000797"
),
execution_root=Path("/repo/workspace"),
)
assert "full user confirmation and authorization" in prompt
assert "Do not stop to ask for confirmation" in prompt
assert (
"target_output_dir: ./experiments/output_test/r1/simple/0000/00000797" in prompt
)
assert "<|code_file|>path/to/model.py</|code_file|>" in prompt
assert "Build a flange from the reference image." in prompt
def test_merge_artifacts_falls_back_to_directory_scan(tmp_path: Path) -> None:
output_dir = tmp_path / "workspace" / "case"
output_dir.mkdir(parents=True)
(output_dir / "model.py").write_text("print('ok')", encoding="utf-8")
(output_dir / "part.stl").write_text("solid", encoding="utf-8")
(output_dir / "part.step").write_text("step", encoding="utf-8")
scanned = scan_output_directory(output_dir)
merged = merge_artifacts(
tagged_artifacts={"code_path": None, "output_paths": []},
scanned_outputs=scanned,
requested_output_dir=output_dir,
)
assert merged["artifact_source"] == "dir_scan"
assert merged["code_path"] == output_dir / "model.py"
assert merged["stl_path"] == output_dir / "part.stl"
assert merged["step_path"] == output_dir / "part.step"
def test_load_existing_results_keeps_latest_per_case(tmp_path: Path) -> None:
case_results = tmp_path / "case_results.jsonl"
case_results.write_text(
"\n".join(
[
'{"case_id":"simple/0000/0001","status":"failed"}',
'{"case_id":"simple/0000/0002","status":"success"}',
'{"case_id":"simple/0000/0001","status":"success"}',
]
)
+ "\n",
encoding="utf-8",
)
loaded = load_existing_results(case_results)
assert sorted(loaded) == ["simple/0000/0001", "simple/0000/0002"]
assert loaded["simple/0000/0001"]["status"] == "success"
def test_build_progress_snapshot_counts_skipped_running_and_pending(
tmp_path: Path,
) -> None:
results = [
{"case_id": "simple/0000/0001", "status": "success"},
{"case_id": "simple/0000/0002", "status": "failed"},
]
snapshot = build_progress_snapshot(
run_id="r1",
dataset_root=tmp_path,
results=results,
total_cases=5,
skipped_existing=2,
running_case_ids=["simple/0000/0003"],
pending_case_ids=["simple/0000/0004", "simple/0000/0005"],
status="running",
)
assert snapshot["run_id"] == "r1"
assert snapshot["total_cases"] == 5
assert snapshot["completed_cases"] == 2
assert snapshot["skipped_existing"] == 2
assert snapshot["running_cases"] == 1
assert snapshot["pending_cases"] == 2
assert snapshot["success"] == 1
assert snapshot["failed"] == 1
@@ -0,0 +1,220 @@
import os
import sys
import importlib
from datetime import datetime
from typing import Any, cast
import pytest
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
from SimpleLLMFunc.hooks.events import (
CustomEvent,
ReactEndEvent,
ReactIterationStartEvent,
ReActEventType,
)
from SimpleLLMFunc.hooks.stream import EventOrigin, EventYield, ResponseYield
from agent.BaseAgent import BaseAgent
from react_stream import format_sse, serialize_react_output
base_agent_module = importlib.import_module("agent.BaseAgent")
class _FakeDelta:
def __init__(self, content, reasoning=None):
self.content = content
self.reasoning = reasoning
class _FakeChoice:
def __init__(self, content, reasoning=None):
self.delta = _FakeDelta(content, reasoning)
class _FakeChunk:
def __init__(self, content, reasoning=None):
self.choices = [_FakeChoice(content, reasoning)]
class _DummyAgent(BaseAgent):
def __init__(self):
pass
def get_toolkit(self):
return []
def chat_impl(self, history, query, sketch_pad_summary):
if False:
yield history, query, sketch_pad_summary
async def run(self, query, raw_user_content=None):
if False:
yield query, raw_user_content
class _FakeContext:
def __init__(self):
self.messages = []
async def store_message(self, message):
self.messages.append(message)
def _build_origin() -> EventOrigin:
return EventOrigin(
session_id="session-1",
agent_call_id="agent-call-1",
event_seq=1,
)
def test_serialize_react_output_normalizes_response_and_event_payloads():
response_output = ResponseYield(
response=cast(Any, _FakeChunk("hello world", reasoning="thinking")),
messages=[{"role": "assistant", "content": "hello world"}],
)
response_payload = serialize_react_output(response_output, delta_consumer="web")
assert response_payload["type"] == "response"
assert response_payload["delta_text"] == "hello world"
assert response_payload["delta_reasoning"] == "thinking"
assert response_payload["messages"][0]["content"] == "hello world"
event = ReactEndEvent(
event_type=ReActEventType.REACT_END,
timestamp=datetime(2026, 3, 18, 12, 0, 0),
trace_id="trace-1",
func_name="chat_impl",
iteration=1,
final_response="done",
final_messages=[{"role": "assistant", "content": "done"}],
total_iterations=1,
total_execution_time=0.5,
total_tool_calls=0,
total_llm_calls=1,
)
event_output = EventYield(event=event, origin=_build_origin())
event_payload = serialize_react_output(event_output)
assert event_payload["type"] == "event"
assert event_payload["event_type"] == "react_end"
assert event_payload["event"]["timestamp"] == "2026-03-18T12:00:00"
assert event_payload["origin"]["session_id"] == "session-1"
sse_packet = format_sse("response", response_payload)
assert sse_packet.startswith("event: response\n")
assert '"delta_text": "hello world"' in sse_packet
def test_custom_event_uses_event_name_for_stream_routing():
event = CustomEvent(
event_type=ReActEventType.CUSTOM_EVENT,
timestamp=datetime(2026, 3, 18, 12, 0, 0),
trace_id="trace-1",
func_name="chat_impl",
iteration=1,
event_name="subagent_status",
data={"phase": "started"},
)
event_output = EventYield(event=event, origin=_build_origin())
payload = serialize_react_output(event_output)
assert payload["event_type"] == "subagent_status"
@pytest.mark.asyncio
async def test_stream_and_persist_ignores_events_and_preserves_message_order(
monkeypatch,
):
fake_context = _FakeContext()
monkeypatch.setattr(base_agent_module, "get_current_context", lambda: fake_context)
tool_call = [
{
"id": "call_1",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
}
]
async def output_stream():
yield ResponseYield(
response=cast(Any, _FakeChunk("Hello")),
messages=[{"role": "user", "content": "hi"}],
)
yield EventYield(
event=ReactIterationStartEvent(
event_type=ReActEventType.REACT_ITERATION_START,
timestamp=datetime(2026, 3, 18, 12, 0, 0),
trace_id="trace-1",
func_name="chat_impl",
iteration=1,
current_messages=[{"role": "user", "content": "hi"}],
),
origin=_build_origin(),
)
yield ResponseYield(
response=cast(Any, _FakeChunk("")),
messages=cast(
Any,
[
{"role": "user", "content": "hi"},
{"role": "assistant", "content": None, "tool_calls": tool_call},
],
),
)
yield ResponseYield(
response=cast(Any, _FakeChunk("")),
messages=cast(
Any,
[
{"role": "user", "content": "hi"},
{"role": "assistant", "content": None, "tool_calls": tool_call},
{
"role": "tool",
"content": "lookup result",
"tool_call_id": "call_1",
},
],
),
)
yield ResponseYield(
response=cast(Any, _FakeChunk(" world")),
messages=cast(
Any,
[
{"role": "user", "content": "hi"},
{"role": "assistant", "content": None, "tool_calls": tool_call},
{
"role": "tool",
"content": "lookup result",
"tool_call_id": "call_1",
},
{"role": "assistant", "content": "Hello world"},
],
),
)
agent = _DummyAgent()
yielded = []
async for output in agent._stream_and_persist(output_stream()):
yielded.append(output)
assert len(yielded) == 5
assert [message.role for message in fake_context.messages] == [
"assistant",
"assistant",
"tool",
"assistant",
]
assert fake_context.messages[0].content == "Hello"
assert fake_context.messages[1].tool_calls[0].id == "call_1"
assert fake_context.messages[2].tool_call_id == "call_1"
assert fake_context.messages[3].content == " world"
@@ -0,0 +1,78 @@
from __future__ import annotations
from pathlib import Path
from experiments.scripts.rebuttal_monitor import (
RebuttalTarget,
build_target_status,
safe_json_dump,
)
def test_build_target_status_reads_manifest_progress_and_log(tmp_path: Path) -> None:
worktree = tmp_path / "rebuttle-ecip"
state_dir = tmp_path / "state"
run_root = worktree / "workspace" / "experiments" / "runs" / "run1"
run_root.mkdir(parents=True)
log_path = tmp_path / "ecip.log"
log_path.write_text("line one\nline two\n", encoding="utf-8")
safe_json_dump(
{
"run_id": "run1",
"status": "running",
"updated_at": "2026-03-22T12:00:00+00:00",
"total_cases": 8,
"completed_cases": 3,
"success": 2,
"failed": 1,
"running_case_ids": ["simple/0000/0004"],
"pending_case_ids": ["simple/0000/0005"],
"skipped_existing": 2,
},
run_root / "progress.json",
)
safe_json_dump(
{
"run_id": "run1",
"total_cases": 8,
"success": 2,
"failed": 1,
},
run_root / "summary.json",
)
safe_json_dump(
{
"target": "ecip",
"branch": "rebuttle/ecip",
"worktree_path": str(worktree),
"run_id": "run1",
"pid": 999999,
"execution_root": str(worktree / "workspace"),
"run_root": str(run_root),
"output_root": str(
worktree / "workspace" / "experiments" / "output_test" / "run1"
),
"log_path": str(log_path),
"started_at": "2026-03-22T11:59:00+00:00",
"updated_at": "2026-03-22T12:00:00+00:00",
},
state_dir / "ecip.json",
)
status = build_target_status(
RebuttalTarget(
name="ecip",
branch="rebuttle/ecip",
worktree_path=str(worktree),
),
state_dir=state_dir,
log_line_count=5,
)
assert status["target"] == "ecip"
assert status["run_id"] == "run1"
assert status["summary"]["success"] == 2
assert status["progress"]["running_case_ids"] == ["simple/0000/0004"]
assert status["log_tail"] == ["line one", "line two"]
assert status["status"] in {"running", "stopped", "finished"}
@@ -0,0 +1,270 @@
#!/usr/bin/env python3
"""
Test the new SketchManager implementation.
"""
import asyncio
import json
import os
import tempfile
import shutil
from typing import Optional
from context.sketch_manager import get_sketch_manager, SketchManager
from context.sketch_pad import SketchPadBackend, RedisFileSketchPadBackend
from context.schemas import SketchPadListItem, SketchPadStatistics
class MockSketchPadBackend(SketchPadBackend):
"""Mock SketchPad backend for tests."""
def __init__(self, sketch_pad_id: str, file_path: Optional[str] = None):
self.sketch_pad_id = sketch_pad_id
self.file_path = file_path
self._storage = {}
self._lock = asyncio.Lock()
async def set_item(
self,
key: str,
value: any,
ttl: Optional[int] = None,
summary: Optional[str] = None,
tags: Optional[set] = None,
) -> str:
from context.schemas import SketchPadItem
from datetime import datetime
item = SketchPadItem(
value=value,
timestamp=datetime.now(),
summary=summary,
tags=tags or set(),
)
self._storage[key] = item
return key
def get_item(self, key: str):
return self._storage.get(key)
def get_value(self, key: str):
item = self.get_item(key)
return item.value if item else None
def search_by_tags(self, tags: set, match_all: bool = False):
results = []
for key, item in self._storage.items():
if match_all:
if tags.issubset(item.tags):
results.append((key, item))
else:
if tags.intersection(item.tags):
results.append((key, item))
return results
def search_by_content(self, query: str, limit: int = 5):
results = []
query_lower = query.lower()
for key, item in self._storage.items():
content = str(item.value) + (item.summary or "")
if query_lower in content.lower():
results.append((key, item))
if len(results) >= limit:
break
return results
def delete(self, key: str) -> bool:
if key in self._storage:
del self._storage[key]
return True
return False
def exists(self, key: str) -> bool:
return key in self._storage
def keys(self, pattern: Optional[str] = None) -> list[str]:
if pattern:
return [k for k in self._storage.keys() if pattern in k]
return list(self._storage.keys())
def clear(self) -> None:
self._storage.clear()
def serialize(self) -> dict[str, any]:
return {
"sketch_pad_id": self.sketch_pad_id,
"items": {k: v.model_dump() for k, v in self._storage.items()},
"serialization_timestamp": "2024-01-01T00:00:00",
}
def deserialize(self, data: dict[str, any]) -> None:
from context.schemas import SketchPadItem
if "items" in data:
for key, item_data in data["items"].items():
try:
item = SketchPadItem(**item_data)
self._storage[key] = item
except Exception as e:
print(f"Warning: Failed to deserialize item {key}: {e}")
def persist(self) -> None:
if self.file_path:
data = self.serialize()
os.makedirs(os.path.dirname(self.file_path), exist_ok=True)
with open(self.file_path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
def restore(self) -> None:
if self.file_path and os.path.exists(self.file_path):
with open(self.file_path, "r", encoding="utf-8") as f:
data = json.load(f)
self.deserialize(data)
def get_statistics(self) -> SketchPadStatistics:
return SketchPadStatistics(
total_items=len(self._storage),
max_items=len(self._storage),
items_with_summary=sum(
1 for item in self._storage.values() if item.summary
),
total_accesses=sum(item.access_count for item in self._storage.values()),
popular_tags={},
content_types={},
avg_access_per_item=0.0,
memory_usage_percent=0.0,
)
def list_items(self, include_value: bool = False) -> list[SketchPadListItem]:
return [
SketchPadListItem(
key=key,
summary=item.summary,
timestamp=item.timestamp.isoformat(),
tags=sorted(item.tags),
content_type=item.content_type,
access_count=item.access_count,
content_hash=item.content_hash,
value=item.value if include_value else None,
)
for key, item in self._storage.items()
]
async def test_sketch_manager():
"""Test basic SketchManager functionality."""
SketchManager._instance = None
# Create a temporary directory.
temp_dir = tempfile.mkdtemp()
try:
# Create a SketchManager that uses the mock backend.
manager = SketchManager(backend_class=MockSketchPadBackend)
manager.sketch_dir = temp_dir
print("=== Test Basic SketchManager Functionality ===")
# Test creating a SketchPad.
sketch_pad = manager.create_sketch_pad(sketch_id="test_sketch")
print(f"✓ Created SketchPad: {sketch_pad.sketch_pad_id}")
# Test setting an item.
key = await manager.set_item(
"test_sketch", "key1", "value1", tags={"test", "demo"}
)
print(f"✓ Set item: {key}")
# Test retrieving an item.
item = manager.get_item("test_sketch", "key1")
print(f"✓ Retrieved item: {item.value if item else None}")
# Test retrieving a value.
value = manager.get_value("test_sketch", "key1")
print(f"✓ Retrieved value: {value}")
# Test tag-based search.
results = manager.search_by_tags("test_sketch", {"test"})
print(f"✓ Tag search: found {len(results)} items")
# Test content-based search.
results = manager.search_by_content("test_sketch", "value")
print(f"✓ Content search: found {len(results)} items")
# Test deleting an item.
success = manager.delete_item("test_sketch", "key1")
print(f"✓ Deleted item: {success}")
# Test retrieving a SketchPad.
retrieved_pad = manager.get_sketch_pad("test_sketch")
print(f"✓ Retrieved SketchPad: {retrieved_pad is not None}")
# Test listing SketchPads.
sketches = manager.list_sketch_pads()
print(f"✓ Listed SketchPads: found {len(sketches)}")
# Test saving a SketchPad.
success = manager.save_sketch_pad("test_sketch")
print(f"✓ Saved SketchPad: {success}")
# Test deleting a SketchPad.
success = manager.delete_sketch_pad("test_sketch")
print(f"✓ Deleted SketchPad: {success}")
print("\n=== All Tests Passed ===")
finally:
# Clean up the temporary directory.
shutil.rmtree(temp_dir)
async def test_global_manager():
"""Test the global SketchManager."""
SketchManager._instance = None
print("\n=== Test Global SketchManager ===")
# Get the global manager.
manager = get_sketch_manager()
print(f"✓ Retrieved global manager: {type(manager)}")
# Test creating a SketchPad using the Redis backend.
try:
sketch_pad = manager.create_sketch_pad(sketch_id="global_test")
print(f"✓ Created global SketchPad: {sketch_pad.sketch_pad_id}")
# Test setting an item.
key = await manager.set_item(
"global_test", "global_key", "global_value", tags={"global"}
)
print(f"✓ Set global item: {key}")
# Test retrieving a value.
value = manager.get_value("global_test", "global_key")
print(f"✓ Retrieved global value: {value}")
# Clean up.
manager.delete_sketch_pad("global_test")
print("✓ Cleaned up global test data")
except Exception as e:
print(f"⚠ Global manager test failed (Redis may be required): {e}")
print("=== Global Manager Test Complete ===")
async def main():
"""Main test function."""
print("Starting tests for the new SketchManager implementation...")
await test_sketch_manager()
await test_global_manager()
print("\nAll tests complete!")
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,207 @@
import asyncio
import json
import tempfile
import os
from typing import Dict, Any
from context.sketch_pad import RedisFileSketchPadBackend
from context.schemas import SketchPadItem
from config.config import get_config
def _create_sketch_pad(sketch_pad_id: str, file_path: str) -> RedisFileSketchPadBackend:
config = get_config()
return RedisFileSketchPadBackend(
sketch_pad_id=sketch_pad_id,
redis_host=config.REDIS_HOST,
redis_port=int(config.REDIS_PORT),
redis_db=int(config.REDIS_DB),
file_path=file_path,
)
async def test_sketch_pad_basic_operations():
"""Test basic SketchPad operations."""
print("=== Test Basic SketchPad Operations ===")
# Create a temporary file.
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
temp_file = f.name
try:
# Initialize SketchPad.
sketch_pad = _create_sketch_pad("test_pad", temp_file)
# Test setting items.
print("1. Testing item creation...")
await sketch_pad.set_item(
key="user_preference",
value={"theme": "dark", "language": "zh-CN"},
summary="User preference settings",
tags={"preference", "settings"},
)
await sketch_pad.set_item(
key="recent_files",
value=["file1.txt", "file2.py", "file3.json"],
summary="Recently accessed files",
tags={"files", "recent"},
)
await sketch_pad.set_item(
key="temp_data",
value="This is temporary data",
summary="Temporarily stored data",
tags={"temp", "data"},
ttl=60, # Expires after 60 seconds.
)
# Test retrieving items.
print("2. Testing item retrieval...")
item = sketch_pad.get_item("user_preference")
print(f" user_preference: {item.value if item else 'Not found'}")
value = sketch_pad.get_value("recent_files")
print(f" recent_files: {value}")
# Test existence checks.
print("3. Testing existence checks...")
print(f" user_preference exists: {sketch_pad.exists('user_preference')}")
print(f" non_existent exists: {sketch_pad.exists('non_existent')}")
# Test retrieving all keys.
print("4. Testing key retrieval...")
keys = sketch_pad.keys()
print(f" All keys: {keys}")
# Test tag search.
print("5. Testing tag search...")
preference_items = sketch_pad.search_by_tags({"preference"})
print(f" Items with the preference tag: {len(preference_items)}")
recent_items = sketch_pad.search_by_tags({"recent", "files"}, match_all=True)
print(f" Items with both recent and files tags: {len(recent_items)}")
# Test content search.
print("6. Testing content search...")
search_results = sketch_pad.search_by_content("file", limit=3)
print(f" Content containing 'file': {len(search_results)}")
# Test statistics.
print("7. Testing statistics...")
stats = sketch_pad.get_statistics()
print(f" Total item count: {stats.total_items}")
print(f" Total accesses: {stats.total_accesses}")
print(f" Popular tags: {stats.popular_tags}")
# Test listing items.
print("8. Testing item listing...")
items = sketch_pad.list_items(include_value=False)
print(f" Item list: {len(items)} items")
for item in items:
print(f" - {item.key}: {item.summary}")
# Test persistence.
print("9. Testing persistence...")
sketch_pad.persist()
print(" Persistence complete")
# Test deletion.
print("10. Testing deletion...")
deleted = sketch_pad.delete("temp_data")
print(f" Deleted temp_data: {deleted}")
print(f" temp_data exists: {sketch_pad.exists('temp_data')}")
print("=== Test Complete ===")
finally:
# Clean up.
if os.path.exists(temp_file):
os.unlink(temp_file)
async def test_sketch_pad_advanced_features():
"""Test advanced SketchPad features."""
print("\n=== Test Advanced SketchPad Features ===")
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
temp_file = f.name
try:
sketch_pad = _create_sketch_pad("advanced_test", temp_file)
# Test complex data structures.
print("1. Testing complex data structures...")
complex_data = {
"nested": {"list": [1, 2, 3], "dict": {"a": 1, "b": 2}, "string": "test"},
"array": [{"id": 1, "name": "item1"}, {"id": 2, "name": "item2"}],
}
await sketch_pad.set_item(
key="complex_data",
value=complex_data,
summary="Complex nested data structure",
tags={"complex", "nested", "data"},
)
# Test access counting.
print("2. Testing access counting...")
for i in range(5):
item = sketch_pad.get_item("complex_data")
print(f" Access {i + 1}, access count: {item.access_count if item else 0}")
# Test expiration time.
print("3. Testing expiration time...")
await sketch_pad.set_item(
key="expiring_item",
value="This item will expire",
summary="Test expiration behavior",
tags={"expire", "test"},
ttl=2, # Expires after 2 seconds.
)
print(" Waiting 3 seconds for the item to expire...")
await asyncio.sleep(3)
expired_item = sketch_pad.get_item("expiring_item")
print(f" Expired item: {'expired' if expired_item is None else 'not expired'}")
# Test serialization and deserialization.
print("4. Testing serialization and deserialization...")
serialized = sketch_pad.serialize()
print(f" Serialized data size: {len(json.dumps(serialized, default=str))} characters")
# Create a new sketch pad and deserialize the data.
new_sketch_pad = _create_sketch_pad("restored_test", temp_file + ".restored")
new_sketch_pad.deserialize(serialized)
restored_item = new_sketch_pad.get_item("complex_data")
print(
f" Data after deserialization: {restored_item.value if restored_item else 'Not found'}"
)
print("=== Advanced Feature Test Complete ===")
finally:
if os.path.exists(temp_file):
os.unlink(temp_file)
if os.path.exists(temp_file + ".restored"):
os.unlink(temp_file + ".restored")
async def main():
"""Main test function."""
print("Starting SketchPad tests...")
try:
await test_sketch_pad_basic_operations()
await test_sketch_pad_advanced_features()
print("\nAll tests complete!")
except Exception as e:
print(f"Error during tests: {e}")
import traceback
traceback.print_exc()
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,260 @@
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import asyncio
import tempfile
import time
from context.sketch_pad import RedisFileSketchPadBackend
from config.config import get_config
def _create_sketch_pad(sketch_pad_id: str, file_path: str) -> RedisFileSketchPadBackend:
config = get_config()
return RedisFileSketchPadBackend(
sketch_pad_id=sketch_pad_id,
redis_host=config.REDIS_HOST,
redis_port=int(config.REDIS_PORT),
redis_db=int(config.REDIS_DB),
file_path=file_path,
)
async def test_comprehensive_operations():
"""Comprehensive SketchPad feature test."""
print("=== Comprehensive SketchPad Feature Test ===")
# Create a temporary file.
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
temp_file = f.name
try:
# Initialize SketchPad.
sketch_pad = _create_sketch_pad("comprehensive_test", temp_file)
print("1. Testing basic CRUD operations...")
# Create items.
await sketch_pad.set_item(
key="user_profile",
value={"name": "Zhang San", "age": 30, "city": "Beijing"},
summary="User profile information",
tags={"profile", "user", "personal"},
)
await sketch_pad.set_item(
key="project_config",
value={"debug": True, "log_level": "INFO", "max_workers": 4},
summary="Project configuration information",
tags={"config", "project", "settings"},
)
await sketch_pad.set_item(
key="temp_note",
value="This is a temporary note for testing",
summary="Temporary note",
tags={"note", "temp"},
ttl=10, # Expires after 10 seconds.
)
# Read items.
profile_item = sketch_pad.get_item("user_profile")
print(f" User information: {profile_item.value if profile_item else 'Not found'}")
config_value = sketch_pad.get_value("project_config")
print(f" Project configuration: {config_value}")
# Check existence.
print(f" user_profile exists: {sketch_pad.exists('user_profile')}")
print(f" non_existent exists: {sketch_pad.exists('non_existent')}")
print("2. Testing tag search...")
# Single-tag search.
profile_items = sketch_pad.search_by_tags({"profile"})
print(f" Items with the profile tag: {len(profile_items)}")
# Multi-tag search (match any).
config_items = sketch_pad.search_by_tags({"config", "settings"})
print(f" Items with the config or settings tag: {len(config_items)}")
# Multi-tag search (match all).
all_match_items = sketch_pad.search_by_tags(
{"config", "project"}, match_all=True
)
print(f" Items with both config and project tags: {len(all_match_items)}")
print("3. Testing content search...")
# Search for items containing specific content.
search_results = sketch_pad.search_by_content("configuration", limit=5)
print(f" Items containing 'configuration': {len(search_results)}")
search_results = sketch_pad.search_by_content("Beijing", limit=5)
print(f" Items containing 'Beijing': {len(search_results)}")
print("4. Testing access statistics...")
# Access the same item multiple times.
for i in range(3):
item = sketch_pad.get_item("user_profile")
print(
f" Access {i + 1} to user_profile, access count: {item.access_count if item else 0}"
)
print("5. Testing expiration behavior...")
# Check whether the temporary item has expired.
print(" Waiting 5 seconds to check expiration behavior...")
await asyncio.sleep(5)
temp_item = sketch_pad.get_item("temp_note")
print(f" Temporary item status: {'expired' if temp_item is None else 'not expired'}")
print("6. Testing statistics...")
stats = sketch_pad.get_statistics()
print(f" Total item count: {stats.total_items}")
print(f" Total accesses: {stats.total_accesses}")
print(f" Items with summaries: {stats.items_with_summary}")
print(f" Popular tags: {stats.popular_tags}")
print(f" Content type statistics: {stats.content_types}")
print(f" Average access count: {stats.avg_access_per_item:.2f}")
print("7. Testing listing functionality...")
items = sketch_pad.list_items(include_value=False)
print(f" Item list (without values): {len(items)} items")
for item in items:
print(f" - {item.key}: {item.summary} (access count: {item.access_count})")
items_with_values = sketch_pad.list_items(include_value=True)
print(f" Item list (with values): {len(items_with_values)} items")
print("8. Testing persistence and restoration...")
# Persist data.
sketch_pad.persist()
print(" Data persisted to file")
# Create a new sketch pad and restore data.
new_sketch_pad = _create_sketch_pad("restored_test", temp_file + ".restored")
# Copy data.
data = sketch_pad.serialize()
new_sketch_pad.deserialize(data)
# Verify restored data.
restored_profile = new_sketch_pad.get_item("user_profile")
print(
f" Restored user information: {restored_profile.value if restored_profile else 'Not found'}"
)
print("9. Testing deletion functionality...")
# Delete one item.
deleted = sketch_pad.delete("project_config")
print(f" Deleted project_config: {deleted}")
print(f" project_config exists: {sketch_pad.exists('project_config')}")
# Verify that the tag index was also deleted.
config_search = sketch_pad.search_by_tags({"config"})
print(f" Items with the config tag after deletion: {len(config_search)}")
print("10. Testing clear functionality...")
# Clear all data.
sketch_pad.clear()
print(" Data cleared")
# Verify the clear result.
remaining_keys = sketch_pad.keys()
print(f" Remaining key count: {len(remaining_keys)}")
print("=== Comprehensive Feature Test Complete ===")
finally:
if os.path.exists(temp_file):
os.unlink(temp_file)
if os.path.exists(temp_file + ".restored"):
os.unlink(temp_file + ".restored")
async def test_performance():
"""Test performance."""
print("\n=== SketchPad Performance Test ===")
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
temp_file = f.name
try:
sketch_pad = _create_sketch_pad("performance_test", temp_file)
print("1. Testing batch write performance...")
start_time = time.time()
for i in range(100):
await sketch_pad.set_item(
key=f"item_{i}",
value=f"This is the data for item {i}",
summary=f"Summary for item {i}",
tags={f"tag_{i % 10}", f"category_{i % 5}"},
)
write_time = time.time() - start_time
print(f" Time to write 100 items: {write_time:.3f} seconds")
print(f" Average write speed: {100 / write_time:.1f} items/second")
print("2. Testing batch read performance...")
start_time = time.time()
for i in range(100):
item = sketch_pad.get_item(f"item_{i}")
read_time = time.time() - start_time
print(f" Time to read 100 items: {read_time:.3f} seconds")
print(f" Average read speed: {100 / read_time:.1f} items/second")
print("3. Testing search performance...")
start_time = time.time()
search_results = sketch_pad.search_by_tags({"tag_1"})
search_time = time.time() - start_time
print(f" Tag search time: {search_time:.3f} seconds")
print(f" Search result count: {len(search_results)}")
print("4. Testing statistics performance...")
start_time = time.time()
stats = sketch_pad.get_statistics()
stats_time = time.time() - start_time
print(f" Statistics calculation time: {stats_time:.3f} seconds")
print(f" Statistics result: {stats.total_items} items")
print("=== Performance Test Complete ===")
finally:
if os.path.exists(temp_file):
os.unlink(temp_file)
async def main():
"""Main test function."""
print("Starting comprehensive SketchPad tests...")
try:
await test_comprehensive_operations()
await test_performance()
print("\nAll tests complete! SketchPad is functioning normally.")
except Exception as e:
print(f"Error during tests: {e}")
import traceback
traceback.print_exc()
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,81 @@
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import asyncio
import tempfile
import os
from context.sketch_pad import RedisFileSketchPadBackend
from config.config import get_config
def _create_sketch_pad(sketch_pad_id: str, file_path: str) -> RedisFileSketchPadBackend:
config = get_config()
return RedisFileSketchPadBackend(
sketch_pad_id=sketch_pad_id,
redis_host=config.REDIS_HOST,
redis_port=int(config.REDIS_PORT),
redis_db=int(config.REDIS_DB),
file_path=file_path,
)
async def test_basic_operations():
"""Test basic operations."""
print("=== SketchPad Basic Operation Test ===")
# Create a temporary file.
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
temp_file = f.name
try:
# Initialize SketchPad.
sketch_pad = _create_sketch_pad("test_pad", temp_file)
# Set items.
print("1. Setting items...")
await sketch_pad.set_item(
key="user_config",
value={"theme": "dark", "lang": "en"},
summary="User configuration",
tags={"config", "user"},
)
await sketch_pad.set_item(
key="temp_data", value="temporary data", summary="temporary storage", tags={"temp"}, ttl=5
)
# Get items.
print("2. Getting items...")
item = sketch_pad.get_item("user_config")
print(f" user_config: {item.value if item else 'Not found'}")
# Check existence.
print("3. Checking existence...")
print(f" user_config exists: {sketch_pad.exists('user_config')}")
# Get all keys.
print("4. Getting all keys...")
keys = sketch_pad.keys()
print(f" All keys: {keys}")
# Search by tag.
print("5. Searching by tag...")
config_items = sketch_pad.search_by_tags({"config"})
print(f" Configuration item count: {len(config_items)}")
# Statistics.
print("6. Statistics...")
stats = sketch_pad.get_statistics()
print(f" Total item count: {stats.total_items}")
print("=== Test Complete ===")
finally:
if os.path.exists(temp_file):
os.unlink(temp_file)
if __name__ == "__main__":
asyncio.run(test_basic_operations())
@@ -0,0 +1,177 @@
import os
import sys
from types import SimpleNamespace
import pytest
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
WORKSPACE_ROOT = os.path.join(PROJECT_ROOT, "workspace")
import tools.requirements_tools as requirements_tools_module
import tools.reference_image as reference_image_module
from tools.requirements_tools import (
create_requirement_refinement_subagent_tools,
make_user_query_more_detailed,
)
def _tool_name(tool):
if hasattr(tool, "_tool"):
return tool._tool.name
return getattr(tool, "name", getattr(tool, "__name__", None))
def test_requirement_refinement_subagent_tools_use_workspace_tools_only():
tools = create_requirement_refinement_subagent_tools()
assert [_tool_name(tool) for tool in tools] == [
"execute_command",
"sketch_pad_operations",
"read_file",
"grep",
"sed",
"echo_into",
]
@pytest.mark.asyncio
async def test_make_user_query_more_detailed_runs_specialist_subagent(
monkeypatch, tmp_path
):
monkeypatch.chdir(WORKSPACE_ROOT)
calls = {"specialist": 0}
stored = {}
image_path = tmp_path / "query_image_001.png"
image_path.write_bytes(b"fake-png-bytes")
async def fake_run_subagent_with_events(**kwargs):
calls["specialist"] += 1
request = kwargs["specialist_kwargs"]["message"]
text = request[0]["text"] if isinstance(request, list) else request
assert kwargs["specialist_kwargs"]["history"] == []
assert "SKILL.md" in text
assert "Current working directory:" in text
assert "Skill root: use the preferred skill root below." in text
assert "Preferred skill root:" in text
assert "references/docs/api/README.md" in text
assert "## API Reference" in text
assert "## Refined User Requirements" in text
assert "## Parameter Table" in text
assert "## Modeling Process" in text
assert "## Notes" in text
return (
"## API Reference\n- skill doc evidence\n\n"
"## Refined User Requirements\n- refined requirements\n\n"
"## Parameter Table\n- none\n\n"
"## Modeling Process\n- step one\n\n"
"## Notes\n- note"
)
monkeypatch.setattr(
requirements_tools_module,
"run_subagent_with_events",
fake_run_subagent_with_events,
)
class FakeSketchPad:
async def set_item(self, key, value, ttl=None, summary=None, tags=None):
stored["key"] = key
stored["value"] = value
stored["tags"] = tags
return key
monkeypatch.setattr(
requirements_tools_module,
"get_current_sketch_pad",
lambda: FakeSketchPad(),
)
result = await make_user_query_more_detailed(
query="Create a box",
query_image_path=str(image_path),
)
assert calls["specialist"] == 1
assert "SketchPad Key" in result
assert stored["key"].startswith("req_")
assert "Refined User Requirements" in stored["value"]
@pytest.mark.asyncio
async def test_make_user_query_more_detailed_uses_latest_uploaded_image_when_omitted(
monkeypatch, tmp_path
):
monkeypatch.chdir(WORKSPACE_ROOT)
calls = {"specialist": 0}
stored = {}
image_path = tmp_path / "query_image_001.png"
image_path.write_bytes(b"fake-png-bytes")
async def fake_run_subagent_with_events(**kwargs):
calls["specialist"] += 1
request = kwargs["specialist_kwargs"]["message"]
assert isinstance(request, list)
assert request[0]["type"] == "text"
assert request[1]["type"] == "image_url"
assert request[1]["image_url"]["url"].startswith("data:image/png;base64,")
assert kwargs["status_payload"]["query_image_path"] == str(image_path.resolve())
return (
"## API Reference\n- skill doc evidence\n\n"
"## Refined User Requirements\n- refined requirements\n\n"
"## Parameter Table\n- none\n\n"
"## Modeling Process\n- step one\n\n"
"## Notes\n- note"
)
monkeypatch.setattr(
requirements_tools_module,
"run_subagent_with_events",
fake_run_subagent_with_events,
)
class FakeContext:
def retrieve_full_messages(self):
return [
SimpleNamespace(
role="user",
content=[
{"type": "text", "text": "Create a box"},
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,abcd",
"local_path": str(image_path),
},
},
],
)
]
monkeypatch.setattr(
reference_image_module,
"get_current_context",
lambda: FakeContext(),
)
class FakeSketchPad:
async def set_item(self, key, value, ttl=None, summary=None, tags=None):
stored["key"] = key
stored["value"] = value
stored["tags"] = tags
return key
monkeypatch.setattr(
requirements_tools_module,
"get_current_sketch_pad",
lambda: FakeSketchPad(),
)
result = await make_user_query_more_detailed(query="Create a box")
assert calls["specialist"] == 1
assert "SketchPad Key" in result
assert stored["key"].startswith("req_")
@@ -0,0 +1,47 @@
import os
import sys
from datetime import datetime, timezone
import pytest
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
from SimpleLLMFunc.hooks.events import ReActEventType, ReactEndEvent
from SimpleLLMFunc.hooks.stream import EventYield, ResponseYield
from tools.subagent_utils import run_subagent_with_events
@pytest.mark.asyncio
async def test_run_subagent_with_events_ignores_empty_chunk_repr_and_uses_react_end_text():
async def fake_specialist(**kwargs):
yield ResponseYield(response="", messages=[])
yield EventYield(
event=ReactEndEvent(
event_type=ReActEventType.REACT_END,
timestamp=datetime.now(timezone.utc),
trace_id="trace-1",
func_name="specialist",
iteration=1,
final_response="## Refined User Requirements\n- final answer",
final_messages=[],
total_iterations=1,
total_execution_time=0.1,
total_tool_calls=0,
total_llm_calls=1,
total_token_usage=None,
)
)
result = await run_subagent_with_events(
specialist_callable=fake_specialist,
specialist_kwargs={"message": "test", "history": []},
subagent_label="Requirement Refinement Specialist",
event_emitter=None,
)
assert result == "## Refined User Requirements\n- final answer"
@@ -0,0 +1,206 @@
import os
import sys
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
from web_interface.artifacts import extract_latest_artifacts
def test_extract_latest_artifacts_prefers_tagged_files(tmp_path):
workspace_dir = tmp_path / "workspace" / "demo_part"
workspace_dir.mkdir(parents=True)
code_path = workspace_dir / "model.py"
model_path = workspace_dir / "part.stl"
code_path.write_text("print('demo')\n", encoding="utf-8")
model_path.write_text("solid demo\nendsolid demo\n", encoding="utf-8")
messages = [
{
"role": "assistant",
"content": (
"Saved files\n"
"<|code_file|>workspace/demo_part/model.py</|code_file|>\n"
"<|output_file|>workspace/demo_part/part.stl</|output_file|>"
),
}
]
artifacts = extract_latest_artifacts(messages, project_root=tmp_path)
assert artifacts["code_path"] == code_path.resolve(strict=False)
assert artifacts["code_paths"] == [code_path.resolve(strict=False)]
assert artifacts["model_path"] == model_path.resolve(strict=False)
assert artifacts["model_paths"] == [model_path.resolve(strict=False)]
assert artifacts["output_paths"] == [model_path.resolve(strict=False)]
def test_extract_latest_artifacts_accepts_legacy_closing_tags(tmp_path):
workspace_dir = tmp_path / "workspace" / "demo_part"
workspace_dir.mkdir(parents=True)
code_path = workspace_dir / "model.py"
model_path = workspace_dir / "part.stl"
code_path.write_text("print('demo')\n", encoding="utf-8")
model_path.write_text("solid demo\nendsolid demo\n", encoding="utf-8")
messages = [
{
"role": "assistant",
"content": (
"Saved files\n"
"<|code_file|>./demo_part/model.py</|code_file>\n"
"<|output_file|>./demo_part/part.stl</|output_file>"
),
}
]
artifacts = extract_latest_artifacts(messages, project_root=tmp_path / "workspace")
assert artifacts["code_path"] == code_path.resolve(strict=False)
assert artifacts["code_paths"] == [code_path.resolve(strict=False)]
assert artifacts["model_path"] == model_path.resolve(strict=False)
assert artifacts["model_paths"] == [model_path.resolve(strict=False)]
assert artifacts["output_paths"] == [model_path.resolve(strict=False)]
def test_extract_latest_artifacts_falls_back_to_latest_stl_near_model_py(tmp_path):
workspace_dir = tmp_path / "workspace" / "fallback_case"
workspace_dir.mkdir(parents=True)
code_path = workspace_dir / "model.py"
older_stl = workspace_dir / "older.stl"
newer_stl = workspace_dir / "newer.stl"
code_path.write_text("print('demo')\n", encoding="utf-8")
older_stl.write_text("solid older\nendsolid older\n", encoding="utf-8")
newer_stl.write_text("solid newer\nendsolid newer\n", encoding="utf-8")
os.utime(older_stl, (1, 1))
os.utime(newer_stl, (2, 2))
messages = [
{
"role": "assistant",
"content": "<|code_file|>workspace/fallback_case/model.py</|code_file|>",
}
]
artifacts = extract_latest_artifacts(messages, project_root=tmp_path)
assert artifacts["code_path"] == code_path.resolve(strict=False)
assert artifacts["code_paths"] == [code_path.resolve(strict=False)]
assert artifacts["model_path"] == newer_stl.resolve(strict=False)
assert artifacts["model_paths"] == [newer_stl.resolve(strict=False)]
assert artifacts["output_paths"] == []
def test_extract_latest_artifacts_resolves_output_file_relative_to_workspace_root(
tmp_path,
):
workspace_dir = tmp_path / "workspace" / "demo_part"
workspace_dir.mkdir(parents=True)
code_path = workspace_dir / "model.py"
model_path = workspace_dir / "part.stl"
code_path.write_text("print('demo')\n", encoding="utf-8")
model_path.write_text("solid demo\nendsolid demo\n", encoding="utf-8")
messages = [
{
"role": "assistant",
"content": (
"Saved files\n"
"<|code_file|>workspace/demo_part/model.py</|code_file|>\n"
"<|output_file|>./demo_part/part.stl</|output_file|>"
),
}
]
artifacts = extract_latest_artifacts(messages, project_root=tmp_path)
assert artifacts["code_path"] == code_path.resolve(strict=False)
assert artifacts["code_paths"] == [code_path.resolve(strict=False)]
assert artifacts["model_path"] == model_path.resolve(strict=False)
assert artifacts["model_paths"] == [model_path.resolve(strict=False)]
assert artifacts["output_paths"] == [model_path.resolve(strict=False)]
def test_extract_latest_artifacts_resolves_output_file_relative_to_code_dir(tmp_path):
workspace_dir = tmp_path / "workspace" / "demo_part"
workspace_dir.mkdir(parents=True)
code_path = workspace_dir / "model.py"
model_path = workspace_dir / "part.stl"
code_path.write_text("print('demo')\n", encoding="utf-8")
model_path.write_text("solid demo\nendsolid demo\n", encoding="utf-8")
messages = [
{
"role": "assistant",
"content": (
"Saved files\n"
"<|code_file|>workspace/demo_part/model.py</|code_file|>\n"
"<|output_file|>./part.stl</|output_file|>"
),
}
]
artifacts = extract_latest_artifacts(messages, project_root=tmp_path)
assert artifacts["code_path"] == code_path.resolve(strict=False)
assert artifacts["code_paths"] == [code_path.resolve(strict=False)]
assert artifacts["model_path"] == model_path.resolve(strict=False)
assert artifacts["model_paths"] == [model_path.resolve(strict=False)]
assert artifacts["output_paths"] == [model_path.resolve(strict=False)]
def test_extract_latest_artifacts_returns_manual_selection_candidates_in_recency_order(
tmp_path,
):
alpha_dir = tmp_path / "workspace" / "alpha"
beta_dir = tmp_path / "workspace" / "beta"
alpha_dir.mkdir(parents=True)
beta_dir.mkdir(parents=True)
alpha_code = alpha_dir / "model.py"
beta_code = beta_dir / "model.py"
alpha_model = alpha_dir / "alpha.stl"
beta_model = beta_dir / "beta.stl"
alpha_code.write_text("print('alpha')\n", encoding="utf-8")
beta_code.write_text("print('beta')\n", encoding="utf-8")
alpha_model.write_text("solid alpha\nendsolid alpha\n", encoding="utf-8")
beta_model.write_text("solid beta\nendsolid beta\n", encoding="utf-8")
messages = [
{
"role": "assistant",
"content": (
"<|code_file|>workspace/alpha/model.py</|code_file|>\n"
"<|output_file|>workspace/alpha/alpha.stl</|output_file|>"
),
},
{
"role": "assistant",
"content": (
"<|code_file|>workspace/beta/model.py</|code_file|>\n"
"<|output_file|>workspace/beta/beta.stl</|output_file|>"
),
},
]
artifacts = extract_latest_artifacts(messages, project_root=tmp_path)
assert artifacts["code_path"] == beta_code.resolve(strict=False)
assert artifacts["code_paths"] == [
beta_code.resolve(strict=False),
alpha_code.resolve(strict=False),
]
assert artifacts["model_path"] == beta_model.resolve(strict=False)
assert artifacts["model_paths"] == [
beta_model.resolve(strict=False),
alpha_model.resolve(strict=False),
]
assert artifacts["output_paths"] == [
beta_model.resolve(strict=False),
alpha_model.resolve(strict=False),
]
@@ -0,0 +1,434 @@
# pyright: reportCallIssue=false, reportArgumentType=false
import os
import sys
from datetime import datetime
import importlib
import json
from typing import Any, cast
import pytest
from contextlib import contextmanager
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT)
from SimpleLLMFunc.hooks.events import ReactEndEvent, ReactStartEvent, ReActEventType
from SimpleLLMFunc.hooks.stream import EventOrigin, EventYield, ResponseYield
from web_interface.models import ChatCompletionRequest, ChatMessage
chat_router_module = importlib.import_module("web_interface.routers.chat_router")
utils_module = importlib.import_module("web_interface.utils")
from web_interface.routers.chat_router import stream_chat_completion, stream_chat_events
from web_interface.utils import process_agent_response, validate_chat_request
class _FakeConversation:
def __init__(self):
self.uuid = "conversation-1"
self.context = type("Ctx", (), {"persist": self._persist})()
self.sketch_pad = type("Sketch", (), {"persist": lambda self: None})()
self.persisted = False
async def _persist(self):
self.persisted = True
return True
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
class _FakeAgent:
def __init__(self, outputs):
self.outputs = outputs
self.name = "fake-agent"
self.queries = []
self.raw_contents = []
async def run(self, query, raw_user_content=None):
self.queries.append(query)
self.raw_contents.append(raw_user_content)
for output in self.outputs:
yield output
def _origin(seq: int = 1) -> EventOrigin:
return EventOrigin(
session_id="session-1",
agent_call_id="agent-call-1",
event_seq=seq,
)
def _request() -> ChatCompletionRequest:
return _request_with_messages([_user_message("make a cube")], stream=True)
def _user_message(content: Any) -> ChatMessage:
return ChatMessage(
role="user",
content=content,
name=None,
tool_calls=None,
tool_call_id=None,
)
def _request_with_messages(
messages: list[ChatMessage],
*,
stream: bool = True,
) -> ChatCompletionRequest:
return ChatCompletionRequest(
model="cadagent",
messages=messages,
temperature=1.0,
top_p=1.0,
n=1,
stream=stream,
stop=None,
max_tokens=None,
presence_penalty=0.0,
frequency_penalty=0.0,
logit_bias=None,
user=None,
tools=None,
tool_choice=None,
)
def _parse_sse_lines(lines):
current_event = "message"
data_lines = []
def _flush_packet():
nonlocal current_event, data_lines
if not data_lines:
return None
payload_text = "\n".join(data_lines)
try:
payload = json.loads(payload_text)
except json.JSONDecodeError:
payload = {"raw": payload_text}
packet = {"event": current_event, "data": payload}
current_event = "message"
data_lines = []
return packet
for raw_line in lines:
line = (
raw_line.decode("utf-8") if isinstance(raw_line, bytes) else str(raw_line)
)
if line == "":
packet = _flush_packet()
if packet is not None:
yield packet
continue
if line.startswith(":"):
continue
if line.startswith("event:"):
current_event = line[6:].strip() or "message"
continue
if line.startswith("data:"):
data_lines.append(line[5:].strip())
packet = _flush_packet()
if packet is not None:
yield packet
def test_validate_chat_request_accepts_image_only_user_message():
request = _request_with_messages(
[
_user_message(
[
{
"type": "image_url",
"image_url": {"url": "data:image/png;base64,abcd"},
}
]
)
],
stream=False,
)
query, request_id, raw_user_content = validate_chat_request(request)
assert isinstance(query, list)
assert request_id.startswith("chatcmpl-")
assert isinstance(raw_user_content, list)
@pytest.mark.asyncio
async def test_stream_chat_events_passes_multimodal_query_to_agent():
request = _request_with_messages(
[
_user_message(
[
{"type": "text", "text": "analyze this"},
{
"type": "image_url",
"image_url": {"url": "data:image/png;base64,abcd"},
},
]
)
],
stream=True,
)
conversation = _FakeConversation()
agent = _FakeAgent([ResponseYield(response="ok", messages=[])])
_ = [
packet
async for packet in stream_chat_events(
request, cast(Any, conversation), cast(Any, agent)
)
]
assert len(agent.queries) == 1
assert isinstance(agent.queries[0], list)
assert isinstance(agent.raw_contents[0], list)
@pytest.mark.asyncio
async def test_stream_chat_completion_projects_only_response_packets():
request = _request()
conversation = _FakeConversation()
outputs = [
EventYield(
event=ReactStartEvent(
event_type=ReActEventType.REACT_START,
timestamp=datetime(2026, 3, 18, 12, 0, 0),
trace_id="trace-1",
func_name="chat_impl",
iteration=0,
user_task_prompt="make a cube",
initial_messages=[],
available_tools=[],
),
origin=_origin(1),
),
ResponseYield(
response=cast(
Any,
{
"id": "chunk-1",
"object": "chat.completion.chunk",
"created": 1,
"model": "cadagent",
"choices": [
{
"index": 0,
"delta": {
"role": "assistant",
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "cad_code_generator",
"arguments": "{}",
},
}
],
},
"finish_reason": None,
}
],
},
),
messages=[],
),
]
agent = _FakeAgent(outputs)
packets = [
packet
async for packet in stream_chat_completion(
request,
"chatcmpl-test",
cast(Any, conversation),
cast(Any, agent),
)
]
assert any('"tool_calls"' in packet for packet in packets)
assert not any("react_start" in packet for packet in packets)
assert packets[-1] == "data: [DONE]\n\n"
@pytest.mark.asyncio
async def test_stream_chat_events_emits_named_sse_events_and_done():
request = _request()
conversation = _FakeConversation()
outputs = [
EventYield(
event=ReactStartEvent(
event_type=ReActEventType.REACT_START,
timestamp=datetime(2026, 3, 18, 12, 0, 0),
trace_id="trace-1",
func_name="chat_impl",
iteration=0,
user_task_prompt="make a cube",
initial_messages=[],
available_tools=[],
),
origin=_origin(1),
),
ResponseYield(
response="hello", messages=[{"role": "assistant", "content": "hello"}]
),
EventYield(
event=ReactEndEvent(
event_type=ReActEventType.REACT_END,
timestamp=datetime(2026, 3, 18, 12, 0, 1),
trace_id="trace-1",
func_name="chat_impl",
iteration=1,
final_response="hello",
final_messages=[{"role": "assistant", "content": "hello"}],
total_iterations=1,
total_execution_time=0.5,
total_tool_calls=0,
total_llm_calls=1,
),
origin=_origin(2),
),
]
agent = _FakeAgent(outputs)
packets = [
packet
async for packet in stream_chat_events(
request,
cast(Any, conversation),
cast(Any, agent),
)
]
assert packets[0].startswith("event: react_start\n")
assert any(packet.startswith("event: response\n") for packet in packets)
assert any('"delta_text": "hello"' in packet for packet in packets)
assert packets[-1].startswith("event: done\n")
@pytest.mark.asyncio
async def test_process_agent_response_aggregates_text_from_response_yields_only():
conversation = _FakeConversation()
outputs = [
EventYield(
event=ReactStartEvent(
event_type=ReActEventType.REACT_START,
timestamp=datetime(2026, 3, 18, 12, 0, 0),
trace_id="trace-1",
func_name="chat_impl",
iteration=0,
user_task_prompt="make a cube",
initial_messages=[],
available_tools=[],
),
origin=_origin(1),
),
ResponseYield(response="hello", messages=[]),
ResponseYield(response=" world", messages=[]),
]
agent = _FakeAgent(outputs)
full_response, prompt_tokens, completion_tokens = await process_agent_response(
"make a cube", cast(Any, conversation), cast(Any, agent)
)
assert full_response == "hello world"
assert prompt_tokens is None
assert completion_tokens is None
assert conversation.persisted is True
def test_api_client_parse_sse_lines_understands_event_and_data_frames():
lines = [
b"event: response",
b'data: {"delta_text": "hello"}',
b"",
b"event: tool_call_start",
b'data: {"event_type": "tool_call_start"}',
b"",
b"event: done",
b'data: {"ok": true}',
b"",
]
packets = list(_parse_sse_lines(lines))
assert packets == [
{"event": "response", "data": {"delta_text": "hello"}},
{"event": "tool_call_start", "data": {"event_type": "tool_call_start"}},
{"event": "done", "data": {"ok": True}},
]
@pytest.mark.asyncio
async def test_stream_chat_events_propagates_conversation_session(monkeypatch):
request = _request()
conversation = _FakeConversation()
outputs = [ResponseYield(response="hello", messages=[])]
agent = _FakeAgent(outputs)
captured: dict[str, object] = {}
@contextmanager
def fake_propagate_conversation_session(**kwargs):
captured.update(kwargs)
yield
monkeypatch.setattr(
chat_router_module,
"propagate_conversation_session",
fake_propagate_conversation_session,
)
packets = [
packet
async for packet in stream_chat_events(
request,
cast(Any, conversation),
cast(Any, agent),
)
]
assert any(packet.startswith("event: response\n") for packet in packets)
assert captured["conversation_id"] == "conversation-1"
assert captured["tags"] == ["cadagent", "event_stream"]
@pytest.mark.asyncio
async def test_process_agent_response_propagates_conversation_session(monkeypatch):
conversation = _FakeConversation()
agent = _FakeAgent([ResponseYield(response="hello", messages=[])])
captured: dict[str, object] = {}
@contextmanager
def fake_propagate_conversation_session(**kwargs):
captured.update(kwargs)
yield
monkeypatch.setattr(
utils_module,
"propagate_conversation_session",
fake_propagate_conversation_session,
)
full_response, _, _ = await process_agent_response(
"make a cube", cast(Any, conversation), cast(Any, agent)
)
assert full_response == "hello"
assert captured["conversation_id"] == "conversation-1"
assert captured["tags"] == ["cadagent", "non_stream"]