first commit
This commit is contained in:
@@ -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.
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user