Compare commits
215
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
59b4abaa14 | ||
|
|
71e65052d1 | ||
|
|
7b0714c5c5 | ||
|
|
4bdcd0b568 | ||
|
|
fdecb76035 | ||
|
|
b7d451ec5d | ||
|
|
86fe3a4749 | ||
|
|
76d5a73cc7 | ||
|
|
2ab6494ec9 | ||
|
|
3f2684dcfe | ||
|
|
266458528e | ||
|
|
35eb35cdc2 | ||
|
|
8cb5d93005 | ||
|
|
5569c99b8e | ||
|
|
d90c3b4a24 | ||
|
|
ee0b25e29a | ||
|
|
a3fe901886 | ||
|
|
153b08f872 | ||
|
|
1b920d7299 | ||
|
|
4b3c42ad5c | ||
|
|
0de186071b | ||
|
|
7bcd6c5349 | ||
|
|
08b399a450 | ||
|
|
97d5bd3c4d | ||
|
|
a8f408b3b0 | ||
|
|
0bdb762832 | ||
|
|
d49e009b12 | ||
|
|
7dc400c05c | ||
|
|
65aca4d260 | ||
|
|
34584c3a2e | ||
|
|
53e09b924c | ||
|
|
1b302ab4bf | ||
|
|
3c681f1639 | ||
|
|
1ff3356d1b | ||
|
|
5193e34803 | ||
|
|
8f8fc81135 | ||
|
|
d4abb3d06f | ||
|
|
b2570f1a62 | ||
|
|
f19b5f5929 | ||
|
|
8e829396b2 | ||
|
|
e8e8ca6700 | ||
|
|
f1cbd4d730 | ||
|
|
f7cebfe7f3 | ||
|
|
b854d9a888 | ||
|
|
83d2acf07f | ||
|
|
eee9c38953 | ||
|
|
e782318338 | ||
|
|
dc94aa76cc | ||
|
|
5cf019c21e | ||
|
|
790bdd6b8a | ||
|
|
b25c09f5ed | ||
|
|
9e8c910ab1 | ||
|
|
cc10e20a47 | ||
|
|
34ed4345fc | ||
|
|
1a85333e4c | ||
|
|
3c587c788a | ||
|
|
303d123527 | ||
|
|
61c2cb4ac4 | ||
|
|
3126b99fdb | ||
|
|
b28b647ce3 | ||
|
|
d736f1cf46 | ||
|
|
e4402f2f83 | ||
|
|
dc5d8edfec | ||
|
|
bdc3be650b | ||
|
|
119de1f347 | ||
|
|
88f6ecea5a | ||
|
|
a0eb6e9dcf | ||
|
|
c987976f82 | ||
|
|
8116848670 | ||
|
|
f6412b8349 | ||
|
|
f1023d9573 | ||
|
|
39560524f7 | ||
|
|
aec8510d49 | ||
|
|
9dd9c0a4be | ||
|
|
53391762be | ||
|
|
d2487ec6a3 | ||
|
|
8e7c6db4b5 | ||
|
|
a16020c4a2 | ||
|
|
83d6e3cd65 | ||
|
|
80e56294f6 | ||
|
|
a1823004aa | ||
|
|
54255c89c4 | ||
|
|
9a7596193f | ||
|
|
31a889f9fd | ||
|
|
6c0d68cfbb | ||
|
|
5234e76c40 | ||
|
|
f048e8cbec | ||
|
|
9e8d2d4a09 | ||
|
|
aa4393b2eb | ||
|
|
9fcb1dc4c0 | ||
|
|
c5f78bf17e | ||
|
|
6003777eda | ||
|
|
bbfb4a0c2c | ||
|
|
b44c94c9cf | ||
|
|
0e16f1bc3d | ||
|
|
83fc393359 | ||
|
|
2e63a09150 | ||
|
|
16791c9717 | ||
|
|
a770d8b0f9 | ||
|
|
f148ffd7a4 | ||
|
|
03d3e3a4da | ||
|
|
6c13a7e722 | ||
|
|
6adefbf190 | ||
|
|
45a377030f | ||
|
|
6165f523c3 | ||
|
|
3c760be0b2 | ||
|
|
b322cd66f2 | ||
|
|
9bb9cd0d55 | ||
|
|
ae1dc44705 | ||
|
|
548179ed3b | ||
|
|
88828112b5 | ||
|
|
e9a221bdd9 | ||
|
|
003c8dbf59 | ||
|
|
51b5e02948 | ||
|
|
97ad118615 | ||
|
|
5d7d526ca0 | ||
|
|
d987d04606 | ||
|
|
e99a47608d | ||
|
|
d584c11b6f | ||
|
|
5de628434c | ||
|
|
fd8365992d | ||
|
|
0fb9504920 | ||
|
|
6e868cb712 | ||
|
|
0ca40f1929 | ||
|
|
195483c65f | ||
|
|
e140b850b4 | ||
|
|
b4c6c4e5ef | ||
|
|
b0e2033ded | ||
|
|
d095bd3cb8 | ||
|
|
5ef45c4345 | ||
|
|
215637a93c | ||
|
|
c6b68f0b6b | ||
|
|
5317bf869b | ||
|
|
0ea1af4ebf | ||
|
|
ca56d15dc9 | ||
|
|
51f38af9eb | ||
|
|
dbd4786b49 | ||
|
|
a2788023a1 | ||
|
|
ba5863f34c | ||
|
|
513582720a | ||
|
|
8e7e94e424 | ||
|
|
31eae748f6 | ||
|
|
ad2d5d2e8f | ||
|
|
a27220dbd0 | ||
|
|
326f18f8a8 | ||
|
|
6f2ff279ae | ||
|
|
41a3366f3e | ||
|
|
5eba972737 | ||
|
|
dafaa3bab4 | ||
|
|
2059acb3a4 | ||
|
|
a8a075600e | ||
|
|
471fd08fba | ||
|
|
1d30c3f6ce | ||
|
|
f959185bca | ||
|
|
1381735e3b | ||
|
|
6612576f8f | ||
|
|
727ffa2943 | ||
|
|
8dc66c713a | ||
|
|
ca8376c4a6 | ||
|
|
e1987c7fa5 | ||
|
|
a267110ce3 | ||
|
|
c00979a3b8 | ||
|
|
171c18eb5a | ||
|
|
5924017c39 | ||
|
|
c34dd1c90f | ||
|
|
891b85403d | ||
|
|
ecf1029b08 | ||
|
|
45294c9cc6 | ||
|
|
c0d7107c07 | ||
|
|
4bec87c1e7 | ||
|
|
422cccf091 | ||
|
|
1579b6e052 | ||
|
|
9e3f4ad5cc | ||
|
|
001cec88c6 | ||
|
|
9854edf113 | ||
|
|
3b08980b9d | ||
|
|
75e840e831 | ||
|
|
95982b62ef | ||
|
|
6bf81753cc | ||
|
|
18482c72e4 | ||
|
|
19a81e1b05 | ||
|
|
73539ef1b8 | ||
|
|
47954fe260 | ||
|
|
38a1c7f838 | ||
|
|
bd03c151da | ||
|
|
8bbe412848 | ||
|
|
d26582c5dc | ||
|
|
b62de7799c | ||
|
|
82ca557b47 | ||
|
|
7fd1017a53 | ||
|
|
c4e78f4d80 | ||
|
|
05f4464935 | ||
|
|
ca9da38a92 | ||
|
|
a541817054 | ||
|
|
37e478a3e2 | ||
|
|
4e94e0a422 | ||
|
|
2d8a5de7b9 | ||
|
|
d6aaff511e | ||
|
|
fe74da53e2 | ||
|
|
dfb81b45a7 | ||
|
|
22629634c9 | ||
|
|
b507630f3b | ||
|
|
68a9b7ad7d | ||
|
|
536ede12de | ||
|
|
6fed0e5e2d | ||
|
|
283a1fbefd | ||
|
|
5b1bc3a47d | ||
|
|
aaf4de23cc | ||
|
|
f639364d7f | ||
|
|
4b0fb7bdbe | ||
|
|
5111c69ff9 | ||
|
|
12e1470506 | ||
|
|
8cb5dd28e6 | ||
|
|
44dc549f75 | ||
|
|
ee2bf70f42 |
+1
-1
@@ -56,7 +56,7 @@ ENV PATH="/root/.local/bin:${PATH}"
|
||||
|
||||
COPY pyproject.toml README.md LICENSE /app/
|
||||
COPY nanobot/ /app/nanobot/
|
||||
RUN uv pip install --system --no-cache --reinstall /app psycopg2-binary
|
||||
RUN uv pip install --system --no-cache --reinstall /app[mem0] psycopg2-binary
|
||||
|
||||
ENTRYPOINT ["nanobot"]
|
||||
CMD ["gateway"]
|
||||
|
||||
@@ -143,7 +143,7 @@ Add or merge these **two parts** into your config (other options have defaults).
|
||||
{
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"model": "anthropic/claude-opus-4-5",
|
||||
"model": "anthropic/claude-opus-4-7",
|
||||
"provider": "openrouter"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
# PR Testing Workflow
|
||||
|
||||
Guide for testing Pull Requests using the local staging environment.
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
./test-pr.sh <pr-number> "test message"
|
||||
```
|
||||
|
||||
## Staging Environment
|
||||
|
||||
**Location:** `/config/workspace/.nanobot-staging/`
|
||||
|
||||
**Components:**
|
||||
- `config.json` — Staging configuration (channels disabled, shared OAuth)
|
||||
- `workspace/` — Isolated workspace for tool operations
|
||||
- `workspace/sessions/` — Session storage (separate from production)
|
||||
|
||||
**Key differences from production:**
|
||||
- No external channels (Telegram disabled)
|
||||
- Uses `NANOBOT_CONFIG` environment variable
|
||||
- Gateway runs on localhost:18791 (vs production's 18790)
|
||||
- `restrictToWorkspace: true` for safety
|
||||
|
||||
## Testing a PR
|
||||
|
||||
### Method 1: Helper Script (Recommended)
|
||||
|
||||
```bash
|
||||
# Test PR with default message
|
||||
./test-pr.sh 31
|
||||
|
||||
# Test with custom message
|
||||
./test-pr.sh 31 "test the hidden message feature"
|
||||
```
|
||||
|
||||
**What it does:**
|
||||
1. Fetches PR from `wylab` remote (force updates if branch exists)
|
||||
2. Checks out PR branch locally
|
||||
3. Installs in editable mode with `uv pip install -e .`
|
||||
4. Runs test with staging config via `NANOBOT_CONFIG` env var
|
||||
5. Leaves branch checked out for further testing
|
||||
|
||||
**After testing:**
|
||||
```bash
|
||||
git checkout main # Return to main branch
|
||||
```
|
||||
|
||||
### Method 2: Manual Testing
|
||||
|
||||
```bash
|
||||
# 1. Fetch and checkout PR
|
||||
cd /config/workspace/nanobot-oauth-port/nanobot-fork
|
||||
git fetch wylab pull/<N>/head:pr-<N>
|
||||
git checkout pr-<N>
|
||||
|
||||
# 2. Install in editable mode
|
||||
uv pip install -e .
|
||||
|
||||
# 3. Test with staging config
|
||||
NANOBOT_CONFIG=/config/workspace/.nanobot-staging/config.json \
|
||||
.venv/bin/nanobot agent -m "test message"
|
||||
|
||||
# 4. For multi-turn testing
|
||||
NANOBOT_CONFIG=/config/workspace/.nanobot-staging/config.json \
|
||||
.venv/bin/nanobot agent # Interactive mode
|
||||
|
||||
# 5. Return to main
|
||||
git checkout main
|
||||
```
|
||||
|
||||
### Method 3: Gateway Validation
|
||||
|
||||
Test that gateway starts without errors:
|
||||
|
||||
```bash
|
||||
NANOBOT_CONFIG=/config/workspace/.nanobot-staging/config.json \
|
||||
.venv/bin/nanobot gateway
|
||||
|
||||
# Kill with Ctrl+C when validated
|
||||
```
|
||||
|
||||
## Verifying Cache Behavior
|
||||
|
||||
To verify prompt caching works correctly (important for performance):
|
||||
|
||||
```bash
|
||||
# Enable logs to see cache metrics
|
||||
NANOBOT_CONFIG=/config/workspace/.nanobot-staging/config.json \
|
||||
.venv/bin/nanobot agent --logs -m "Turn 1: list files"
|
||||
|
||||
# Look for cache metrics in output:
|
||||
# - cache_write: New cache entries created
|
||||
# - cache_read: Tokens read from cache
|
||||
```
|
||||
|
||||
**What to look for:**
|
||||
- Turn 1: High `cache_write`, moderate `cache_read`
|
||||
- Turn 2+: Low `cache_write`, high `cache_read` (reusing cache)
|
||||
- `cache_read` should increase across turns as context grows
|
||||
|
||||
**Example healthy pattern:**
|
||||
```
|
||||
Turn 1: cache_write=354 cache_read=3563
|
||||
Turn 2: cache_write=255 cache_read=3917 ← Same as Turn 1 end
|
||||
Turn 3: cache_write=113 cache_read=4172 ← Growing with context
|
||||
```
|
||||
|
||||
## Session Management
|
||||
|
||||
### Clear session for fresh test
|
||||
|
||||
```bash
|
||||
rm -f /config/workspace/.nanobot-staging/workspace/sessions/cli_direct.jsonl
|
||||
```
|
||||
|
||||
### View session contents
|
||||
|
||||
```bash
|
||||
cat /config/workspace/.nanobot-staging/workspace/sessions/cli_direct.jsonl | jq
|
||||
```
|
||||
|
||||
### Check for specific features (e.g., hidden signatures)
|
||||
|
||||
```bash
|
||||
cat /config/workspace/.nanobot-staging/workspace/sessions/cli_direct.jsonl | grep "_hidden_sig"
|
||||
```
|
||||
|
||||
## Common Testing Scenarios
|
||||
|
||||
### Test tool execution
|
||||
|
||||
```bash
|
||||
./test-pr.sh 31 "List all Python files in the current directory"
|
||||
```
|
||||
|
||||
### Test multi-turn conversation
|
||||
|
||||
```bash
|
||||
NANOBOT_CONFIG=/config/workspace/.nanobot-staging/config.json \
|
||||
.venv/bin/nanobot agent
|
||||
|
||||
# Then interact naturally:
|
||||
> list files in current directory
|
||||
> how many python files are there?
|
||||
> what's the total size?
|
||||
```
|
||||
|
||||
### Test error handling
|
||||
|
||||
```bash
|
||||
./test-pr.sh 31 "Try to read a file that doesn't exist: /nonexistent.txt"
|
||||
```
|
||||
|
||||
### Test with thinking mode
|
||||
|
||||
The staging config has `thinking_budget: 10000` enabled by default, so all tests use extended thinking.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### "No API key configured" error
|
||||
|
||||
- **Cause:** `NANOBOT_CONFIG` env var not set
|
||||
- **Fix:** Ensure you're using `NANOBOT_CONFIG=/config/workspace/.nanobot-staging/config.json`
|
||||
|
||||
### "Module not found" after checkout
|
||||
|
||||
- **Cause:** Need to reinstall after switching branches
|
||||
- **Fix:** Run `uv pip install -e .` after checkout
|
||||
|
||||
### Changes not applying
|
||||
|
||||
- **Cause:** Using cached `.pyc` files
|
||||
- **Fix:** Clear pycache: `find . -type d -name __pycache__ -exec rm -rf {} + 2>/dev/null || true`
|
||||
|
||||
### Session has stale data
|
||||
|
||||
- **Cause:** Previous test left session data
|
||||
- **Fix:** `rm /config/workspace/.nanobot-staging/workspace/sessions/cli_direct.jsonl`
|
||||
|
||||
## Best Practices
|
||||
|
||||
1. **Clear session between PR tests** to avoid cross-contamination
|
||||
2. **Test with tool use** to trigger agentic behavior (not just simple Q&A)
|
||||
3. **Check cache metrics** for performance-sensitive PRs
|
||||
4. **Run with `--logs`** to see detailed behavior during development
|
||||
5. **Return to main** after testing to avoid accidental commits on PR branches
|
||||
|
||||
## Integration with CI/CD
|
||||
|
||||
The staging environment is currently manual-only. Future enhancements:
|
||||
|
||||
- [ ] Automated PR testing via Gitea Actions
|
||||
- [ ] Cache validation in CI pipeline
|
||||
- [ ] Multi-PR parallel testing using git worktrees
|
||||
- [ ] Regression test suite against production behavior
|
||||
|
||||
## File Locations Reference
|
||||
|
||||
| Path | Purpose |
|
||||
|------|---------|
|
||||
| `/config/workspace/nanobot-oauth-port/nanobot-fork/` | Local nanobot repository |
|
||||
| `/config/workspace/.nanobot-staging/` | Staging environment root |
|
||||
| `/config/workspace/.nanobot-staging/config.json` | Staging configuration |
|
||||
| `/config/workspace/.nanobot-staging/workspace/` | Staging workspace |
|
||||
| `/config/workspace/.nanobot-staging/workspace/sessions/` | Session storage |
|
||||
| `/config/workspace/nanobot-oauth-port/nanobot-fork/test-pr.sh` | Helper script |
|
||||
|
||||
## Related Documentation
|
||||
|
||||
- [nanobot README](../README.md) - Main project documentation
|
||||
- [CLAUDE.md](../CLAUDE.md) - Development guide for Claude Code
|
||||
- [config/schema.py](../nanobot/config/schema.py) - Configuration schema
|
||||
@@ -6,8 +6,12 @@ import platform
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.memory import MemoryStore
|
||||
from nanobot.agent.memory_mem0 import Mem0MemoryStore, HAS_MEM0
|
||||
from nanobot.agent.skills import SkillsLoader
|
||||
from nanobot.agent.visibility import compute_signature
|
||||
|
||||
|
||||
class ContextBuilder:
|
||||
@@ -19,10 +23,21 @@ class ContextBuilder:
|
||||
"""
|
||||
|
||||
BOOTSTRAP_FILES = ["AGENTS.md", "SOUL.md", "USER.md", "TOOLS.md", "IDENTITY.md"]
|
||||
|
||||
def __init__(self, workspace: Path):
|
||||
|
||||
def __init__(self, workspace: Path, mem0_config: dict[str, Any] | None = None):
|
||||
self.workspace = workspace
|
||||
self.memory = MemoryStore(workspace)
|
||||
|
||||
# Choose memory backend based on config
|
||||
if mem0_config and mem0_config.get("enabled") and HAS_MEM0:
|
||||
self.memory = Mem0MemoryStore(workspace, config=mem0_config)
|
||||
self.use_mem0 = True
|
||||
logger.info("ContextBuilder using mem0 for semantic memory")
|
||||
else:
|
||||
if mem0_config and mem0_config.get("enabled"):
|
||||
logger.warning("mem0 enabled but not installed, falling back to MEMORY.md")
|
||||
self.memory = MemoryStore(workspace)
|
||||
self.use_mem0 = False
|
||||
|
||||
self.skills = SkillsLoader(workspace)
|
||||
|
||||
def build_system_prompt(self, skill_names: list[str] | None = None) -> str:
|
||||
@@ -152,6 +167,18 @@ visibility markers will be rejected."""
|
||||
system_prompt = self.build_system_prompt(skill_names)
|
||||
if channel and chat_id:
|
||||
system_prompt += f"\n\n## Current Session\nChannel: {channel}\nChat ID: {chat_id}"
|
||||
|
||||
# Add mem0 semantic memory context (if enabled)
|
||||
if self.use_mem0 and channel and chat_id:
|
||||
user_id = f"{channel}_{chat_id}"
|
||||
memory_context = self.memory.get_memory_context(
|
||||
query=current_message,
|
||||
user_id=user_id,
|
||||
limit=5
|
||||
)
|
||||
if memory_context:
|
||||
system_prompt += f"\n\n{memory_context}"
|
||||
|
||||
messages.append({"role": "system", "content": system_prompt})
|
||||
|
||||
# History
|
||||
@@ -200,12 +227,14 @@ visibility markers will be rejected."""
|
||||
Returns:
|
||||
Updated message list.
|
||||
"""
|
||||
messages.append({
|
||||
msg: dict[str, Any] = {
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_call_id,
|
||||
"name": tool_name,
|
||||
"content": result
|
||||
})
|
||||
"content": result,
|
||||
"_hidden_sig": compute_signature(result if isinstance(result, str) else ""),
|
||||
}
|
||||
messages.append(msg)
|
||||
return messages
|
||||
|
||||
def add_assistant_message(
|
||||
@@ -228,13 +257,14 @@ visibility markers will be rejected."""
|
||||
Updated message list.
|
||||
"""
|
||||
msg: dict[str, Any] = {"role": "assistant", "content": content or ""}
|
||||
|
||||
|
||||
if tool_calls:
|
||||
msg["tool_calls"] = tool_calls
|
||||
|
||||
msg["_hidden_sig"] = compute_signature(content or "")
|
||||
|
||||
# Thinking models reject history without this
|
||||
if reasoning_content:
|
||||
msg["reasoning_content"] = reasoning_content
|
||||
|
||||
|
||||
messages.append(msg)
|
||||
return messages
|
||||
|
||||
+926
-414
File diff suppressed because it is too large
Load Diff
@@ -87,11 +87,7 @@ class MemoryStore:
|
||||
keep_count = memory_window // 2
|
||||
if len(session.messages) <= keep_count:
|
||||
return True
|
||||
if len(session.messages) - session.last_consolidated <= 0:
|
||||
return True
|
||||
old_messages = session.messages[session.last_consolidated:-keep_count]
|
||||
if not old_messages:
|
||||
return True
|
||||
old_messages = session.messages[:-keep_count]
|
||||
logger.info("Memory consolidation: {} to consolidate, {} keep", len(old_messages), keep_count)
|
||||
|
||||
lines = []
|
||||
@@ -142,8 +138,7 @@ class MemoryStore:
|
||||
if update != current_memory:
|
||||
self.write_long_term(update)
|
||||
|
||||
session.last_consolidated = 0 if archive_all else len(session.messages) - keep_count
|
||||
logger.info("Memory consolidation done: {} messages, last_consolidated={}", len(session.messages), session.last_consolidated)
|
||||
logger.info("Memory consolidation done: {} messages total", len(session.messages))
|
||||
return True
|
||||
except Exception:
|
||||
logger.exception("Memory consolidation failed")
|
||||
|
||||
@@ -0,0 +1,404 @@
|
||||
"""Mem0-powered memory system for intelligent semantic retrieval."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.providers.base import LLMProvider
|
||||
from nanobot.session.manager import Session
|
||||
|
||||
try:
|
||||
from mem0 import Memory
|
||||
from mem0.configs.base import MemoryConfig
|
||||
HAS_MEM0 = True
|
||||
except ImportError:
|
||||
HAS_MEM0 = False
|
||||
MemoryConfig = None # type: ignore
|
||||
|
||||
|
||||
class Mem0MemoryStore:
|
||||
"""
|
||||
Enhanced memory store using mem0 for semantic search and automatic extraction.
|
||||
|
||||
Features:
|
||||
- Multi-level memory (user, session, agent)
|
||||
- Semantic search with embeddings
|
||||
- Automatic memory extraction from conversations
|
||||
- 90% token reduction vs full-context
|
||||
- 91% faster responses
|
||||
"""
|
||||
|
||||
def __init__(self, workspace: Path, config: dict[str, Any] | None = None):
|
||||
if not HAS_MEM0:
|
||||
raise ImportError(
|
||||
"mem0 not installed. Install with: pip install mem0ai"
|
||||
)
|
||||
|
||||
self.workspace = workspace
|
||||
self.memory_dir = workspace / "memory"
|
||||
self.memory_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Build custom extraction prompt tuned for nanobot conversations
|
||||
from datetime import datetime
|
||||
|
||||
today = datetime.now().strftime("%Y-%m-%d")
|
||||
custom_prompt = f"Extract dated facts from this conversation as JSON: {{\"facts\": [...]}}. Today is {today}.\n\n"
|
||||
self.custom_prompt = custom_prompt
|
||||
|
||||
# Initialize mem0 with optional config + custom prompt
|
||||
# Extract only MemoryConfig-relevant fields
|
||||
raw_config = config if config else {}
|
||||
logger.debug(f"Mem0MemoryStore received config keys: {list(raw_config.keys())}")
|
||||
mem0_cfg_dict = {}
|
||||
for key in ("vector_store", "llm", "embedder", "graph_store", "version"):
|
||||
if key in raw_config:
|
||||
mem0_cfg_dict[key] = raw_config[key]
|
||||
logger.debug(f"Extracted for MemoryConfig: {list(mem0_cfg_dict.keys())}")
|
||||
mem0_config = MemoryConfig(**mem0_cfg_dict)
|
||||
logger.debug(f"MemoryConfig created: vector_store={mem0_config.vector_store.provider if mem0_config.vector_store else None}")
|
||||
self.memory = Memory(config=mem0_config)
|
||||
|
||||
logger.info("Mem0 memory system initialized with custom nanobot prompt")
|
||||
|
||||
def search_memories(
|
||||
self,
|
||||
query: str,
|
||||
user_id: str,
|
||||
limit: int = 5,
|
||||
session_id: str | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Search for relevant memories using semantic search.
|
||||
|
||||
Args:
|
||||
query: Search query (user's current message)
|
||||
user_id: User identifier (e.g., "telegram_12345")
|
||||
limit: Max number of memories to return
|
||||
session_id: Optional session-specific memories
|
||||
|
||||
Returns:
|
||||
List of memory dicts with 'memory' and 'score' keys
|
||||
"""
|
||||
try:
|
||||
# Search user-level memories
|
||||
user_memories = self.memory.search(
|
||||
query=query,
|
||||
user_id=user_id,
|
||||
limit=limit
|
||||
)
|
||||
|
||||
results = []
|
||||
if user_memories and "results" in user_memories:
|
||||
results.extend(user_memories["results"])
|
||||
|
||||
# Optionally search session-level memories
|
||||
if session_id:
|
||||
session_memories = self.memory.search(
|
||||
query=query,
|
||||
user_id=user_id,
|
||||
metadata={"session_id": session_id},
|
||||
limit=limit // 2 # Reserve half for session context
|
||||
)
|
||||
if session_memories and "results" in session_memories:
|
||||
results.extend(session_memories["results"])
|
||||
|
||||
logger.debug(
|
||||
f"Mem0 search: query='{query[:50]}...', found {len(results)} memories"
|
||||
)
|
||||
return results[:limit] # Limit total results
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Mem0 search failed: {e}")
|
||||
return []
|
||||
|
||||
def add_conversation(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
user_id: str,
|
||||
session_id: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Add conversation messages to memory for automatic extraction.
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'
|
||||
user_id: User identifier
|
||||
session_id: Optional session identifier for session-level memories
|
||||
"""
|
||||
try:
|
||||
metadata = {}
|
||||
if session_id:
|
||||
metadata["session_id"] = session_id
|
||||
|
||||
# mem0 automatically extracts and stores relevant facts
|
||||
result = self.memory.add(
|
||||
messages,
|
||||
user_id=user_id,
|
||||
metadata=metadata if metadata else None
|
||||
)
|
||||
|
||||
facts_count = len(result.get("results", [])) if result else 0
|
||||
logger.debug(
|
||||
f"Mem0 add: {len(messages)} messages for user {user_id}, extracted {facts_count} facts"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Mem0 add failed: {e}")
|
||||
|
||||
async def extract_facts(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
provider: Any,
|
||||
model: str,
|
||||
) -> list[str]:
|
||||
"""Extract facts from conversation using the main agent's LLM provider."""
|
||||
import json as _json
|
||||
|
||||
conv_text = ""
|
||||
for msg in messages:
|
||||
role = msg.get("role", "unknown")
|
||||
content_val = msg.get("content", "")
|
||||
if isinstance(content_val, str) and content_val.strip():
|
||||
conv_text += f"{role}: {content_val}\n\n"
|
||||
|
||||
if not conv_text.strip():
|
||||
return []
|
||||
|
||||
extraction_messages = [
|
||||
{"role": "user", "content": self.custom_prompt + conv_text}
|
||||
]
|
||||
|
||||
try:
|
||||
response = await provider.chat(
|
||||
messages=extraction_messages,
|
||||
model=model,
|
||||
max_tokens=16384,
|
||||
temperature=0.3,
|
||||
thinking_budget=0,
|
||||
)
|
||||
text = (response.content or "").strip()
|
||||
if text.startswith("```"):
|
||||
text = text.split("```")[1]
|
||||
if text.startswith("json"):
|
||||
text = text[4:]
|
||||
text = text.strip()
|
||||
data = _json.loads(text)
|
||||
facts = data.get("facts", [])
|
||||
if not isinstance(facts, list):
|
||||
logger.warning(f"LLM returned non-list facts: {type(facts)}")
|
||||
return []
|
||||
logger.debug(f"Extracted {len(facts)} facts using {model}")
|
||||
return facts
|
||||
except Exception as e:
|
||||
logger.error(f"Fact extraction failed: {e}")
|
||||
return []
|
||||
|
||||
def store_facts(
|
||||
self,
|
||||
facts: list[str],
|
||||
user_id: str,
|
||||
session_id: str | None = None,
|
||||
) -> None:
|
||||
"""Store pre-extracted facts in mem0 with infer=False."""
|
||||
if not facts:
|
||||
return
|
||||
|
||||
metadata = {}
|
||||
if session_id:
|
||||
metadata["session_id"] = session_id
|
||||
|
||||
stored = 0
|
||||
for fact in facts:
|
||||
# Normalize: LLM may return dicts like {"fact": "...", "date": "..."} or plain strings
|
||||
if isinstance(fact, dict):
|
||||
fact_text = fact.get("fact", fact.get("text", str(fact)))
|
||||
else:
|
||||
fact_text = str(fact)
|
||||
if not fact_text.strip():
|
||||
continue
|
||||
try:
|
||||
self.memory.add(
|
||||
fact_text,
|
||||
user_id=user_id,
|
||||
infer=False,
|
||||
metadata=metadata if metadata else None,
|
||||
)
|
||||
stored += 1
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to store fact '{str(fact_text)[:50]}...': {e}")
|
||||
|
||||
logger.info(f"Stored {stored}/{len(facts)} facts for user {user_id}")
|
||||
|
||||
def get_memory_context(
|
||||
self,
|
||||
query: str,
|
||||
user_id: str,
|
||||
limit: int = 5
|
||||
) -> str:
|
||||
"""
|
||||
Get formatted memory context for inclusion in system prompt.
|
||||
|
||||
Args:
|
||||
query: Current user query
|
||||
user_id: User identifier
|
||||
limit: Max memories to include
|
||||
|
||||
Returns:
|
||||
Formatted memory context string
|
||||
"""
|
||||
memories = self.search_memories(query, user_id, limit=limit)
|
||||
|
||||
if not memories:
|
||||
return ""
|
||||
|
||||
lines = ["## Relevant Memories"]
|
||||
for i, mem in enumerate(memories, 1):
|
||||
memory_text = mem.get("memory", "")
|
||||
# Include score if available for debugging
|
||||
score = mem.get("score", "")
|
||||
score_str = f" (relevance: {score:.2f})" if score else ""
|
||||
lines.append(f"{i}. {memory_text}{score_str}")
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
def update_memory(self, memory_id: str, data: dict[str, Any]) -> None:
|
||||
"""Update a specific memory by ID."""
|
||||
try:
|
||||
self.memory.update(memory_id, data)
|
||||
logger.debug(f"Mem0 update: memory_id={memory_id}")
|
||||
except Exception as e:
|
||||
logger.error(f"Mem0 update failed: {e}")
|
||||
|
||||
def delete_memory(self, memory_id: str) -> None:
|
||||
"""Delete a specific memory by ID."""
|
||||
try:
|
||||
self.memory.delete(memory_id)
|
||||
logger.debug(f"Mem0 delete: memory_id={memory_id}")
|
||||
except Exception as e:
|
||||
logger.error(f"Mem0 delete failed: {e}")
|
||||
|
||||
def get_all_memories(self, user_id: str) -> list[dict[str, Any]]:
|
||||
"""Get all memories for a user."""
|
||||
try:
|
||||
result = self.memory.get_all(user_id=user_id)
|
||||
return result.get("results", []) if result else []
|
||||
except Exception as e:
|
||||
logger.error(f"Mem0 get_all failed: {e}")
|
||||
return []
|
||||
|
||||
async def consolidate(
|
||||
self,
|
||||
session: Session,
|
||||
provider: LLMProvider,
|
||||
model: str,
|
||||
*,
|
||||
archive_all: bool = False,
|
||||
memory_window: int = 50,
|
||||
) -> bool:
|
||||
"""
|
||||
Consolidate session messages into mem0 memory.
|
||||
|
||||
Unlike the original MemoryStore, mem0 handles extraction automatically,
|
||||
so this just needs to feed recent messages to mem0.
|
||||
|
||||
Returns True on success.
|
||||
"""
|
||||
try:
|
||||
# Extract user_id from session key (e.g., "telegram:12345" -> "telegram_12345")
|
||||
user_id = session.key.replace(":", "_")
|
||||
|
||||
# Determine which messages to consolidate
|
||||
if archive_all:
|
||||
messages_to_add = session.messages
|
||||
logger.info(
|
||||
f"Mem0 consolidation (archive_all): {len(messages_to_add)} messages"
|
||||
)
|
||||
else:
|
||||
keep_count = memory_window // 2
|
||||
if len(session.messages) <= keep_count:
|
||||
return True
|
||||
|
||||
# Consolidate messages except the most recent (kept for context)
|
||||
start_idx = 0
|
||||
end_idx = len(session.messages) - keep_count
|
||||
|
||||
if end_idx <= start_idx:
|
||||
return True
|
||||
|
||||
messages_to_add = session.messages[start_idx:end_idx]
|
||||
|
||||
if not messages_to_add:
|
||||
return True
|
||||
|
||||
logger.info(
|
||||
f"Mem0 consolidation: {len(messages_to_add)} to consolidate, "
|
||||
f"{keep_count} keep"
|
||||
)
|
||||
|
||||
# Convert to mem0 format with intelligent filtering
|
||||
mem0_messages = []
|
||||
for msg in messages_to_add:
|
||||
role = msg.get("role")
|
||||
content = msg.get("content")
|
||||
|
||||
# Skip tool results — raw bash output, file contents, and JSON
|
||||
# get misinterpreted by the extraction LLM as user interests
|
||||
if role == "tool":
|
||||
continue
|
||||
|
||||
# Skip system messages — they're boilerplate instructions, not facts
|
||||
if role == "system":
|
||||
continue
|
||||
|
||||
# Skip messages with no content
|
||||
if not content:
|
||||
continue
|
||||
|
||||
# Normalize assistant message content: extract text from Anthropic list format
|
||||
if role == "assistant" and isinstance(content, list):
|
||||
# Anthropic format: list of {type: "text"|"tool_use", text: "..."} blocks
|
||||
text_parts = [
|
||||
block.get("text", "")
|
||||
for block in content
|
||||
if isinstance(block, dict) and block.get("type") == "text"
|
||||
]
|
||||
content = " ".join(text_parts).strip()
|
||||
if not content:
|
||||
continue # Skip if assistant only called tools with no text explanation
|
||||
|
||||
# Normalize user message content (could also be a list in some formats)
|
||||
if isinstance(content, list):
|
||||
text_parts = [
|
||||
block.get("text", "") if isinstance(block, dict) else str(block)
|
||||
for block in content
|
||||
]
|
||||
content = " ".join(text_parts).strip()
|
||||
if not content:
|
||||
continue
|
||||
|
||||
# Skip trivially short messages (commands like "/new")
|
||||
if len(content.strip()) < 10:
|
||||
continue
|
||||
|
||||
mem0_messages.append({
|
||||
"role": role,
|
||||
"content": content
|
||||
})
|
||||
|
||||
if mem0_messages:
|
||||
# Extract facts using the main agent's LLM (already paid for),
|
||||
# then store with infer=False to bypass mem0's GPT-nano
|
||||
facts = await self.extract_facts(mem0_messages, provider, model)
|
||||
self.store_facts(facts, user_id=user_id, session_id=session.key)
|
||||
|
||||
logger.info(
|
||||
f"Mem0 consolidation done: {len(session.messages)} messages total"
|
||||
)
|
||||
return True
|
||||
|
||||
except Exception:
|
||||
logger.exception("Mem0 consolidation failed")
|
||||
return False
|
||||
@@ -167,10 +167,10 @@ class SkillsLoader:
|
||||
return content
|
||||
|
||||
def _parse_nanobot_metadata(self, raw: str) -> dict:
|
||||
"""Parse skill metadata JSON from frontmatter (supports nanobot and openclaw keys)."""
|
||||
"""Parse skill metadata JSON from frontmatter (supports nanobot, clawdbot, and openclaw keys)."""
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
return (data.get("nanobot") or data.get("openclaw") or data.get("clawdbot") or {}) if isinstance(data, dict) else {}
|
||||
return (data.get("nanobot") or data.get("clawdbot") or data.get("openclaw") or {}) if isinstance(data, dict) else {}
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return {}
|
||||
|
||||
|
||||
@@ -73,7 +73,7 @@ class SubagentManager:
|
||||
origin_metadata: Optional metadata to propagate to announcement (e.g. suppress_output).
|
||||
|
||||
Returns:
|
||||
Status message indicating the subagent was started.
|
||||
Task ID of the spawned subagent.
|
||||
"""
|
||||
task_id = str(uuid.uuid4())[:8]
|
||||
display_label = label or task[:30] + ("..." if len(task) > 30 else "")
|
||||
@@ -83,18 +83,18 @@ class SubagentManager:
|
||||
"chat_id": origin_chat_id,
|
||||
"metadata": origin_metadata or {},
|
||||
}
|
||||
|
||||
|
||||
# Create background task
|
||||
bg_task = asyncio.create_task(
|
||||
self._run_subagent(task_id, task, display_label, origin, model=model)
|
||||
)
|
||||
self._running_tasks[task_id] = bg_task
|
||||
|
||||
|
||||
# Cleanup when done
|
||||
bg_task.add_done_callback(lambda _: self._running_tasks.pop(task_id, None))
|
||||
|
||||
|
||||
logger.info(f"Spawned subagent [{task_id}]: {display_label}")
|
||||
return f"Subagent [{display_label}] started. Task ID: {task_id}"
|
||||
return task_id
|
||||
|
||||
async def _run_subagent(
|
||||
self,
|
||||
|
||||
@@ -1,114 +1,117 @@
|
||||
"""BashTool20250124 - Persistent bash session with sentinel-based output.
|
||||
"""BashTool20250124 - Persistent bash session with async buffer polling.
|
||||
|
||||
Anthropic's native bash_20250124 tool with a long-running session.
|
||||
Based on Anthropic's reference implementation from anthropic-quickstarts.
|
||||
Uses asyncio.create_subprocess_shell + direct buffer reads instead of
|
||||
threaded readline, which avoids exhausting the default ThreadPoolExecutor.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import subprocess
|
||||
import uuid
|
||||
import os
|
||||
from typing import Any, Literal
|
||||
|
||||
from nanobot.agent.tools.anthropic.base import BaseAnthropicTool, ToolResult
|
||||
from nanobot.agent.tools.anthropic.base import BaseAnthropicTool, ToolResult, ToolError
|
||||
|
||||
|
||||
class _BashSession:
|
||||
"""Manages a persistent bash subprocess with sentinel-based output reading."""
|
||||
"""A session of a bash shell.
|
||||
|
||||
Uses asyncio subprocess with direct buffer polling — no threads.
|
||||
Based on anthropics/anthropic-quickstarts computer-use-demo.
|
||||
"""
|
||||
|
||||
command: str = "/bin/bash"
|
||||
_output_delay: float = 0.2 # seconds between buffer polls
|
||||
_timeout: float = 120.0 # seconds
|
||||
_sentinel: str = "<<exit>>"
|
||||
|
||||
def __init__(self):
|
||||
self.process: subprocess.Popen | None = None
|
||||
self._start()
|
||||
self._started = False
|
||||
self._timed_out = False
|
||||
self._process: asyncio.subprocess.Process | None = None
|
||||
|
||||
def _start(self):
|
||||
"""Start the bash process."""
|
||||
self.process = subprocess.Popen(
|
||||
["bash"],
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
bufsize=1,
|
||||
async def start(self):
|
||||
if self._started:
|
||||
return
|
||||
|
||||
self._process = await asyncio.create_subprocess_shell(
|
||||
self.command,
|
||||
preexec_fn=os.setsid,
|
||||
shell=True,
|
||||
bufsize=0,
|
||||
stdin=asyncio.subprocess.PIPE,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
)
|
||||
self._started = True
|
||||
|
||||
def restart(self):
|
||||
"""Restart the bash session."""
|
||||
if self.process:
|
||||
self.process.terminate()
|
||||
try:
|
||||
self.process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
self.process.kill()
|
||||
self.process.wait()
|
||||
self._start()
|
||||
def stop(self):
|
||||
"""Terminate the bash shell."""
|
||||
if not self._started:
|
||||
return
|
||||
if self._process and self._process.returncode is None:
|
||||
self._process.terminate()
|
||||
|
||||
async def run_command(self, command: str, timeout: float = 120.0) -> str:
|
||||
"""Run a command in the persistent bash session.
|
||||
async def run(self, command: str) -> ToolResult:
|
||||
"""Execute a command in the bash shell."""
|
||||
if not self._started:
|
||||
raise ToolError("Session has not started.")
|
||||
if self._process is None or self._process.returncode is not None:
|
||||
return ToolResult(
|
||||
system="tool must be restarted",
|
||||
error=f"bash has exited with returncode "
|
||||
f"{self._process.returncode if self._process else 'unknown'}",
|
||||
)
|
||||
if self._timed_out:
|
||||
raise ToolError(
|
||||
f"timed out: bash has not returned in {self._timeout} seconds "
|
||||
"and must be restarted",
|
||||
)
|
||||
|
||||
Uses a unique sentinel to detect command completion.
|
||||
assert self._process.stdin
|
||||
assert self._process.stdout
|
||||
assert self._process.stderr
|
||||
|
||||
Args:
|
||||
command: Bash command to execute
|
||||
timeout: Maximum time to wait for command completion (seconds)
|
||||
# Send command + sentinel on its own line so heredoc terminators
|
||||
# aren't corrupted (EOF; echo '...' ≠ EOF)
|
||||
self._process.stdin.write(
|
||||
command.encode() + f"\necho '{self._sentinel}'\n".encode()
|
||||
)
|
||||
await self._process.stdin.drain()
|
||||
|
||||
Returns:
|
||||
Command output (stdout + stderr combined)
|
||||
# Poll stdout buffer until sentinel appears — no threads involved
|
||||
try:
|
||||
async with asyncio.timeout(self._timeout):
|
||||
while True:
|
||||
await asyncio.sleep(self._output_delay)
|
||||
output = self._process.stdout._buffer.decode()
|
||||
if self._sentinel in output:
|
||||
output = output[: output.index(self._sentinel)]
|
||||
break
|
||||
except asyncio.TimeoutError:
|
||||
self._timed_out = True
|
||||
raise ToolError(
|
||||
f"timed out: bash has not returned in {self._timeout} seconds "
|
||||
"and must be restarted",
|
||||
) from None
|
||||
|
||||
Raises:
|
||||
asyncio.TimeoutError: If command doesn't complete within timeout
|
||||
RuntimeError: If bash process has died
|
||||
"""
|
||||
if not self.process or self.process.poll() is not None:
|
||||
raise RuntimeError("Bash process has died")
|
||||
if output.endswith("\n"):
|
||||
output = output[:-1]
|
||||
|
||||
# Generate unique sentinel
|
||||
sentinel = f"<<BASH_COMMAND_DONE_{uuid.uuid4().hex}>>"
|
||||
error = self._process.stderr._buffer.decode()
|
||||
if error.endswith("\n"):
|
||||
error = error[:-1]
|
||||
|
||||
# Send command + sentinel
|
||||
full_command = f"{command}\necho '{sentinel}'\n"
|
||||
self.process.stdin.write(full_command)
|
||||
self.process.stdin.flush()
|
||||
# Clear buffers for next command
|
||||
self._process.stdout._buffer.clear()
|
||||
self._process.stderr._buffer.clear()
|
||||
|
||||
# Read output until sentinel appears
|
||||
output_lines = []
|
||||
start_time = asyncio.get_event_loop().time()
|
||||
|
||||
while True:
|
||||
# Check timeout
|
||||
elapsed = asyncio.get_event_loop().time() - start_time
|
||||
if elapsed > timeout:
|
||||
raise asyncio.TimeoutError(
|
||||
f"Command timed out after {timeout}s: {command[:50]}..."
|
||||
)
|
||||
|
||||
# Read line (non-blocking via asyncio)
|
||||
try:
|
||||
line = await asyncio.wait_for(
|
||||
asyncio.to_thread(self.process.stdout.readline),
|
||||
timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
# No output yet, continue waiting
|
||||
continue
|
||||
|
||||
if not line:
|
||||
# EOF - process died
|
||||
raise RuntimeError("Bash process terminated unexpectedly")
|
||||
|
||||
# Check for sentinel
|
||||
if sentinel in line:
|
||||
break
|
||||
|
||||
output_lines.append(line.rstrip("\n"))
|
||||
|
||||
return "\n".join(output_lines)
|
||||
|
||||
def __del__(self):
|
||||
"""Clean up bash process on deletion."""
|
||||
if self.process:
|
||||
self.process.terminate()
|
||||
try:
|
||||
self.process.wait(timeout=2)
|
||||
except subprocess.TimeoutExpired:
|
||||
self.process.kill()
|
||||
# Return as ToolResult (our loop handles this type)
|
||||
if error and output:
|
||||
return ToolResult(output=f"{output}\n\nstderr: {error}")
|
||||
elif error:
|
||||
return ToolResult(output=error)
|
||||
else:
|
||||
return ToolResult(output=output if output else "(no output)")
|
||||
|
||||
|
||||
class BashTool20250124(BaseAnthropicTool):
|
||||
@@ -124,10 +127,10 @@ class BashTool20250124(BaseAnthropicTool):
|
||||
|
||||
api_type: Literal["bash_20250124"] = "bash_20250124"
|
||||
name: Literal["bash"] = "bash"
|
||||
beta_flag: str = "computer-use-2025-11-24"
|
||||
beta_flag: str | None = None
|
||||
|
||||
def __init__(self):
|
||||
self._session = _BashSession()
|
||||
self._session: _BashSession | None = None
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
@@ -135,39 +138,26 @@ class BashTool20250124(BaseAnthropicTool):
|
||||
restart: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> ToolResult:
|
||||
"""Execute bash command or restart session.
|
||||
|
||||
Args:
|
||||
command: Bash command to execute (optional)
|
||||
restart: Restart the bash session (optional)
|
||||
**kwargs: Additional arguments (ignored)
|
||||
|
||||
Returns:
|
||||
ToolResult with command output or error
|
||||
"""
|
||||
if restart:
|
||||
self._session.restart()
|
||||
return ToolResult(output="Bash session restarted successfully.")
|
||||
if self._session:
|
||||
self._session.stop()
|
||||
self._session = _BashSession()
|
||||
await self._session.start()
|
||||
return ToolResult(system="tool has been restarted.")
|
||||
|
||||
if not command:
|
||||
return ToolResult(
|
||||
error="Either 'command' or 'restart=True' must be provided."
|
||||
)
|
||||
if self._session is None:
|
||||
self._session = _BashSession()
|
||||
await self._session.start()
|
||||
|
||||
try:
|
||||
output = await self._session.run_command(command)
|
||||
return ToolResult(output=output if output else "(no output)")
|
||||
except asyncio.TimeoutError as e:
|
||||
return ToolResult(error=f"Command timed out: {e}")
|
||||
except Exception as e:
|
||||
return ToolResult(error=f"{e}")
|
||||
if command is not None:
|
||||
try:
|
||||
return await self._session.run(command)
|
||||
except ToolError as e:
|
||||
return ToolResult(error=str(e))
|
||||
|
||||
return ToolResult(error="Either 'command' or 'restart=True' must be provided.")
|
||||
|
||||
def to_params(self) -> dict[str, Any]:
|
||||
"""Convert to Anthropic API tool parameter format.
|
||||
|
||||
Returns:
|
||||
Tool definition for Anthropic API with bash_20250124 type
|
||||
"""
|
||||
return {
|
||||
"type": self.api_type,
|
||||
"name": self.name,
|
||||
|
||||
@@ -67,13 +67,14 @@ class ComputerTool20251124(BaseAnthropicTool):
|
||||
self.display_height_px = display_height_px
|
||||
|
||||
def to_params(self):
|
||||
"""Return tool definition for API."""
|
||||
"""Return tool definition for API.
|
||||
|
||||
NOTE: display_width_px, display_height_px, and enable_zoom are NOT
|
||||
valid parameters for computer_20251124 and cause API hangs if sent.
|
||||
"""
|
||||
return {
|
||||
"type": self.api_type,
|
||||
"name": self.name,
|
||||
"display_width_px": self.display_width_px,
|
||||
"display_height_px": self.display_height_px,
|
||||
"enable_zoom": True,
|
||||
}
|
||||
|
||||
async def __call__(
|
||||
|
||||
@@ -20,7 +20,7 @@ class EditTool20250728(BaseAnthropicTool):
|
||||
|
||||
api_type: Literal["text_editor_20250728"] = "text_editor_20250728"
|
||||
name: Literal["str_replace_based_edit_tool"] = "str_replace_based_edit_tool"
|
||||
beta_flag: str = "computer-use-2025-11-24"
|
||||
beta_flag: str | None = None
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
"""Mem0 memory tools — expose semantic memory to the agent."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, TYPE_CHECKING
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.tools.base import Tool
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.memory_mem0 import Mem0MemoryStore
|
||||
|
||||
|
||||
class Mem0ToolContext:
|
||||
"""Shared mutable state injected into every mem0 tool."""
|
||||
|
||||
def __init__(self, store: Mem0MemoryStore, consolidate_fn):
|
||||
self.store = store
|
||||
self.consolidate_fn = consolidate_fn # async (session, archive_all) -> None
|
||||
self.user_id: str = "unknown"
|
||||
self.session = None
|
||||
|
||||
def set_context(self, channel: str, chat_id: str, session=None):
|
||||
self.user_id = f"{channel}_{chat_id}"
|
||||
self.session = session
|
||||
|
||||
|
||||
class MemorySearchTool(Tool):
|
||||
"""Search memories semantically."""
|
||||
|
||||
name = "memory_search"
|
||||
description = (
|
||||
"Search your long-term memory for facts relevant to a query. "
|
||||
"Returns the most relevant memories ranked by similarity."
|
||||
)
|
||||
parameters = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {
|
||||
"type": "string",
|
||||
"description": "Natural-language search query",
|
||||
},
|
||||
"limit": {
|
||||
"type": "integer",
|
||||
"description": "Max results to return (default 5)",
|
||||
"minimum": 1,
|
||||
"maximum": 20,
|
||||
},
|
||||
},
|
||||
"required": ["query"],
|
||||
}
|
||||
|
||||
def __init__(self, ctx: Mem0ToolContext):
|
||||
self._ctx = ctx
|
||||
|
||||
async def execute(self, query: str, limit: int = 5, **kw: Any) -> str:
|
||||
results = self._ctx.store.search_memories(
|
||||
query=query,
|
||||
user_id=self._ctx.user_id,
|
||||
limit=limit,
|
||||
)
|
||||
if not results:
|
||||
return "No memories found."
|
||||
lines = []
|
||||
for i, mem in enumerate(results, 1):
|
||||
text = mem.get("memory", "")
|
||||
score = mem.get("score")
|
||||
mid = mem.get("id", "")
|
||||
score_str = f" (score: {score:.2f})" if score else ""
|
||||
lines.append(f"{i}. [{mid}] {text}{score_str}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
class MemoryListTool(Tool):
|
||||
"""List all memories for the current user."""
|
||||
|
||||
name = "memory_list"
|
||||
description = (
|
||||
"List ALL stored memories for the current user. "
|
||||
"Use memory_search for targeted lookup; use this to browse everything."
|
||||
)
|
||||
parameters = {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
}
|
||||
|
||||
def __init__(self, ctx: Mem0ToolContext):
|
||||
self._ctx = ctx
|
||||
|
||||
async def execute(self, **kw: Any) -> str:
|
||||
memories = self._ctx.store.get_all_memories(self._ctx.user_id)
|
||||
if not memories:
|
||||
return "No memories stored."
|
||||
lines = []
|
||||
for i, mem in enumerate(memories, 1):
|
||||
text = mem.get("memory", "")
|
||||
mid = mem.get("id", "")
|
||||
lines.append(f"{i}. [{mid}] {text}")
|
||||
return f"{len(memories)} memories:\n" + "\n".join(lines)
|
||||
|
||||
|
||||
class MemoryAddTool(Tool):
|
||||
"""Add a fact to long-term memory."""
|
||||
|
||||
name = "memory_add"
|
||||
description = (
|
||||
"Store a new fact or piece of information in long-term memory. "
|
||||
"The content will be processed by the extraction LLM and stored as one or more facts."
|
||||
)
|
||||
parameters = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"content": {
|
||||
"type": "string",
|
||||
"description": "The fact or information to remember",
|
||||
},
|
||||
},
|
||||
"required": ["content"],
|
||||
}
|
||||
|
||||
def __init__(self, ctx: Mem0ToolContext):
|
||||
self._ctx = ctx
|
||||
|
||||
async def execute(self, content: str, **kw: Any) -> str:
|
||||
try:
|
||||
result = self._ctx.store.memory.add(
|
||||
[{"role": "user", "content": content}],
|
||||
user_id=self._ctx.user_id,
|
||||
)
|
||||
facts_count = len(result.get("results", [])) if result else 0
|
||||
return f"Added to memory. {facts_count} fact(s) extracted."
|
||||
except Exception as e:
|
||||
logger.error(f"memory_add failed: {e}")
|
||||
return f"Error adding memory: {e}"
|
||||
|
||||
|
||||
class MemoryUpdateTool(Tool):
|
||||
"""Update an existing memory by ID."""
|
||||
|
||||
name = "memory_update"
|
||||
description = (
|
||||
"Update the content of an existing memory. "
|
||||
"Use memory_list or memory_search first to find the memory ID."
|
||||
)
|
||||
parameters = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"memory_id": {
|
||||
"type": "string",
|
||||
"description": "The memory ID to update",
|
||||
},
|
||||
"content": {
|
||||
"type": "string",
|
||||
"description": "The new content for this memory",
|
||||
},
|
||||
},
|
||||
"required": ["memory_id", "content"],
|
||||
}
|
||||
|
||||
def __init__(self, ctx: Mem0ToolContext):
|
||||
self._ctx = ctx
|
||||
|
||||
async def execute(self, memory_id: str, content: str, **kw: Any) -> str:
|
||||
try:
|
||||
self._ctx.store.update_memory(memory_id, content)
|
||||
return f"Memory {memory_id} updated."
|
||||
except Exception as e:
|
||||
logger.error(f"memory_update failed: {e}")
|
||||
return f"Error updating memory: {e}"
|
||||
|
||||
|
||||
class MemoryDeleteTool(Tool):
|
||||
"""Delete a memory by ID."""
|
||||
|
||||
name = "memory_delete"
|
||||
description = (
|
||||
"Delete a specific memory by its ID. "
|
||||
"Use memory_list or memory_search first to find the memory ID."
|
||||
)
|
||||
parameters = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"memory_id": {
|
||||
"type": "string",
|
||||
"description": "The memory ID to delete",
|
||||
},
|
||||
},
|
||||
"required": ["memory_id"],
|
||||
}
|
||||
|
||||
def __init__(self, ctx: Mem0ToolContext):
|
||||
self._ctx = ctx
|
||||
|
||||
async def execute(self, memory_id: str, **kw: Any) -> str:
|
||||
try:
|
||||
self._ctx.store.delete_memory(memory_id)
|
||||
return f"Memory {memory_id} deleted."
|
||||
except Exception as e:
|
||||
logger.error(f"memory_delete failed: {e}")
|
||||
return f"Error deleting memory: {e}"
|
||||
|
||||
|
||||
class MemoryConsolidateTool(Tool):
|
||||
"""Trigger memory consolidation for the current session."""
|
||||
|
||||
name = "memory_consolidate"
|
||||
description = (
|
||||
"Extract and store facts from the current conversation into long-term memory. "
|
||||
"Normally this happens automatically on /new, but you can trigger it manually."
|
||||
)
|
||||
parameters = {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
}
|
||||
|
||||
def __init__(self, ctx: Mem0ToolContext):
|
||||
self._ctx = ctx
|
||||
|
||||
async def execute(self, **kw: Any) -> str:
|
||||
session = self._ctx.session
|
||||
if not session:
|
||||
return "Error: no active session."
|
||||
try:
|
||||
await self._ctx.consolidate_fn(session, archive_all=False)
|
||||
return "Memory consolidation complete."
|
||||
except Exception as e:
|
||||
logger.error(f"memory_consolidate failed: {e}")
|
||||
return f"Error during consolidation: {e}"
|
||||
@@ -21,6 +21,7 @@ class MessageTool(Tool):
|
||||
self._sessions = sessions
|
||||
self._default_channel = default_channel
|
||||
self._default_chat_id = default_chat_id
|
||||
self._sent_in_turn: bool = False
|
||||
|
||||
def set_context(self, channel: str, chat_id: str) -> None:
|
||||
"""Set the current message context."""
|
||||
@@ -30,6 +31,10 @@ class MessageTool(Tool):
|
||||
def set_send_callback(self, callback: Callable[[OutboundMessage], Awaitable[None]]) -> None:
|
||||
"""Set the callback for sending messages."""
|
||||
self._send_callback = callback
|
||||
|
||||
def start_turn(self) -> None:
|
||||
"""Reset per-turn send tracking."""
|
||||
self._sent_in_turn = False
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
@@ -92,6 +97,10 @@ class MessageTool(Tool):
|
||||
try:
|
||||
await self._send_callback(msg)
|
||||
|
||||
# Track if sent to same target as current context
|
||||
if channel == self._default_channel and chat_id == self._default_chat_id:
|
||||
self._sent_in_turn = True
|
||||
|
||||
if self._sessions:
|
||||
session_key = f"{channel}:{chat_id}"
|
||||
session = self._sessions.get_or_create(session_key)
|
||||
|
||||
+12
-13
@@ -4,11 +4,19 @@
|
||||
import hmac
|
||||
import hashlib
|
||||
import re
|
||||
from typing import Tuple
|
||||
|
||||
SECRET_KEY = "nanobot_visibility_secret_key_v1"
|
||||
|
||||
|
||||
def compute_signature(content: str) -> str:
|
||||
"""Compute HMAC signature for content (hex string, no prefix)."""
|
||||
return hmac.new(
|
||||
SECRET_KEY.encode(),
|
||||
content.encode(),
|
||||
hashlib.sha256
|
||||
).hexdigest()[:8]
|
||||
|
||||
|
||||
def sign_content(content: str) -> str:
|
||||
"""
|
||||
Sign content with HMAC and prepend marker.
|
||||
@@ -19,15 +27,11 @@ def sign_content(content: str) -> str:
|
||||
Returns:
|
||||
Content with signed visibility marker: "[HIDDEN:{sig}] {content}"
|
||||
"""
|
||||
sig = hmac.new(
|
||||
SECRET_KEY.encode(),
|
||||
content.encode(),
|
||||
hashlib.sha256
|
||||
).hexdigest()[:8]
|
||||
sig = compute_signature(content)
|
||||
return f"[HIDDEN:{sig}] {content}"
|
||||
|
||||
|
||||
def verify_signature(marked_content: str) -> Tuple[bool, str]:
|
||||
def verify_signature(marked_content: str) -> tuple[bool, str]:
|
||||
"""
|
||||
Verify HMAC signature and extract clean content.
|
||||
|
||||
@@ -44,12 +48,7 @@ def verify_signature(marked_content: str) -> Tuple[bool, str]:
|
||||
return False, marked_content
|
||||
|
||||
claimed_sig, content = match.groups()
|
||||
expected_sig = hmac.new(
|
||||
SECRET_KEY.encode(),
|
||||
content.encode(),
|
||||
hashlib.sha256
|
||||
).hexdigest()[:8]
|
||||
|
||||
expected_sig = compute_signature(content)
|
||||
is_valid = hmac.compare_digest(claimed_sig, expected_sig)
|
||||
return is_valid, content
|
||||
|
||||
|
||||
+9
-11
@@ -1,9 +1,7 @@
|
||||
"""Async message queue for decoupled channel-agent communication."""
|
||||
|
||||
import asyncio
|
||||
from typing import Callable, Awaitable
|
||||
|
||||
from loguru import logger
|
||||
from typing import Awaitable, Callable
|
||||
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
|
||||
@@ -11,30 +9,30 @@ from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
class MessageBus:
|
||||
"""
|
||||
Async message bus that decouples chat channels from the agent core.
|
||||
|
||||
|
||||
Channels push messages to the inbound queue, and the agent processes
|
||||
them and pushes responses to the outbound queue.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self):
|
||||
self.inbound: asyncio.Queue[InboundMessage] = asyncio.Queue()
|
||||
self.outbound: asyncio.Queue[OutboundMessage] = asyncio.Queue()
|
||||
self._outbound_subscribers: dict[str, list[Callable[[OutboundMessage], Awaitable[None]]]] = {}
|
||||
self._correlation_store: dict[str, asyncio.Future] = {}
|
||||
self._running = False
|
||||
|
||||
|
||||
async def publish_inbound(self, msg: InboundMessage) -> None:
|
||||
"""Publish a message from a channel to the agent."""
|
||||
await self.inbound.put(msg)
|
||||
|
||||
|
||||
async def consume_inbound(self) -> InboundMessage:
|
||||
"""Consume the next inbound message (blocks until available)."""
|
||||
return await self.inbound.get()
|
||||
|
||||
|
||||
async def publish_outbound(self, msg: OutboundMessage) -> None:
|
||||
"""Publish a response from the agent to channels."""
|
||||
await self.outbound.put(msg)
|
||||
|
||||
|
||||
async def consume_outbound(self) -> OutboundMessage:
|
||||
"""Consume the next outbound message (blocks until available)."""
|
||||
return await self.outbound.get()
|
||||
@@ -91,12 +89,12 @@ class MessageBus:
|
||||
def stop(self) -> None:
|
||||
"""Stop the dispatcher loop."""
|
||||
self._running = False
|
||||
|
||||
|
||||
@property
|
||||
def inbound_size(self) -> int:
|
||||
"""Number of pending inbound messages."""
|
||||
return self.inbound.qsize()
|
||||
|
||||
|
||||
@property
|
||||
def outbound_size(self) -> int:
|
||||
"""Number of pending outbound messages."""
|
||||
|
||||
@@ -210,15 +210,15 @@ class ChannelManager:
|
||||
timeout=1.0
|
||||
)
|
||||
|
||||
# Resolve any pending correlation (hook request-response)
|
||||
self.bus.resolve_correlation(msg)
|
||||
|
||||
if msg.metadata.get("_progress"):
|
||||
if msg.metadata.get("_tool_hint") and not self.config.channels.send_tool_hints:
|
||||
continue
|
||||
if not msg.metadata.get("_tool_hint") and not self.config.channels.send_progress:
|
||||
continue
|
||||
|
||||
# Resolve any pending correlation (hook request-response)
|
||||
self.bus.resolve_correlation(msg)
|
||||
|
||||
channel = self.channels.get(msg.channel)
|
||||
if channel:
|
||||
try:
|
||||
|
||||
+309
-410
File diff suppressed because it is too large
Load Diff
@@ -1,13 +1,21 @@
|
||||
"""Configuration loading utilities."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from nanobot.config.schema import Config
|
||||
|
||||
|
||||
def get_config_path() -> Path:
|
||||
"""Get the default configuration file path."""
|
||||
"""Get the configuration file path.
|
||||
|
||||
Checks NANOBOT_CONFIG environment variable first, otherwise defaults
|
||||
to ~/.nanobot/config.json
|
||||
"""
|
||||
env_path = os.getenv("NANOBOT_CONFIG")
|
||||
if env_path:
|
||||
return Path(env_path)
|
||||
return Path.home() / ".nanobot" / "config.json"
|
||||
|
||||
|
||||
@@ -84,4 +92,18 @@ def _migrate_config(data: dict) -> dict:
|
||||
exec_cfg = tools.get("exec", {})
|
||||
if "restrictToWorkspace" in exec_cfg and "restrictToWorkspace" not in tools:
|
||||
tools["restrictToWorkspace"] = exec_cfg.pop("restrictToWorkspace")
|
||||
|
||||
# Extract api_key from oauthCredentials if present
|
||||
providers = data.get("providers", {})
|
||||
for _, provider_config in providers.items():
|
||||
if isinstance(provider_config, dict):
|
||||
oauth_creds = provider_config.get("oauthCredentials")
|
||||
if oauth_creds and isinstance(oauth_creds, dict):
|
||||
access_token = oauth_creds.get("access_token", "")
|
||||
# Only set api_key if not already set and access_token exists
|
||||
if access_token and not provider_config.get("api_key"):
|
||||
provider_config["api_key"] = access_token
|
||||
# Clean up migrated data to avoid duplication
|
||||
del provider_config["oauthCredentials"]
|
||||
|
||||
return data
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Configuration schema using Pydantic."""
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field, ConfigDict
|
||||
from pydantic.alias_generators import to_camel
|
||||
@@ -220,7 +220,7 @@ class AgentDefaults(Base):
|
||||
"""Default agent configuration."""
|
||||
|
||||
workspace: str = "~/.nanobot/workspace"
|
||||
model: str = "anthropic/claude-opus-4-5"
|
||||
model: str = "anthropic/claude-opus-4-7"
|
||||
provider: str = "auto" # Provider name (e.g. "anthropic", "openrouter") or "auto" for auto-detection
|
||||
max_tokens: int = 8192
|
||||
temperature: float = 0.1
|
||||
@@ -310,7 +310,7 @@ class GatewayConfig(Base):
|
||||
heartbeat: HeartbeatConfig = Field(default_factory=HeartbeatConfig)
|
||||
|
||||
|
||||
class HooksConfig(Base):
|
||||
class HooksConfig(BaseModel):
|
||||
"""Webhook endpoint configuration."""
|
||||
enabled: bool = False
|
||||
tokens: dict[str, str] = Field(default_factory=dict) # Named tokens: {name: secret}
|
||||
@@ -330,7 +330,7 @@ class HooksConfig(Base):
|
||||
return bool(self.tokens)
|
||||
|
||||
|
||||
class WebSearchConfig(Base):
|
||||
class WebSearchConfig(BaseModel):
|
||||
"""Web search tool configuration."""
|
||||
|
||||
api_key: str = "" # Brave Search API key
|
||||
@@ -361,6 +361,17 @@ class MCPServerConfig(Base):
|
||||
tool_timeout: int = 30 # Seconds before a tool call is cancelled
|
||||
|
||||
|
||||
class Mem0Config(Base):
|
||||
"""Mem0 memory system configuration."""
|
||||
|
||||
enabled: bool = False # If true, use mem0 for semantic memory instead of simple MEMORY.md
|
||||
api_key: str = "" # Optional: mem0 cloud API key (leave empty for self-hosted)
|
||||
search_limit: int = 5 # Max memories to retrieve per query
|
||||
llm: str = "" # Optional: LLM for memory extraction (default: gpt-4.1-nano-2025-04-14)
|
||||
embedder: str = "" # Optional: Embedding model (default: mem0's default)
|
||||
vector_store: dict[str, Any] = Field(default_factory=dict) # Vector store config (e.g., {"provider": "qdrant", "config": {...}})
|
||||
|
||||
|
||||
class ToolsConfig(Base):
|
||||
"""Tools configuration."""
|
||||
|
||||
@@ -368,6 +379,7 @@ class ToolsConfig(Base):
|
||||
exec: ExecToolConfig = Field(default_factory=ExecToolConfig)
|
||||
restrict_to_workspace: bool = False # If true, restrict all tool access to workspace directory
|
||||
enable_memory_tool: bool = True # If true, enable Anthropic's native memory tool
|
||||
mem0: Mem0Config = Field(default_factory=Mem0Config) # Mem0 semantic memory configuration
|
||||
mcp_servers: dict[str, MCPServerConfig] = Field(default_factory=dict)
|
||||
|
||||
|
||||
|
||||
@@ -87,6 +87,10 @@ class HeartbeatService:
|
||||
logger.info("Heartbeat disabled")
|
||||
return
|
||||
|
||||
# Idempotent: don't create a new task if already running
|
||||
if self._task is not None and not self._task.done():
|
||||
return
|
||||
|
||||
self._running = True
|
||||
self._task = asyncio.create_task(self._run_loop())
|
||||
logger.info(f"Heartbeat started (every {self.interval_s}s)")
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""Provider module exports."""
|
||||
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse, LongContextError, ToolCallRequest
|
||||
from nanobot.providers.litellm_provider import LiteLLMProvider
|
||||
from nanobot.providers.openai_codex_provider import OpenAICodexProvider
|
||||
from nanobot.providers.anthropic_oauth import AnthropicOAuthProvider
|
||||
|
||||
@@ -10,8 +10,8 @@ from typing import Any
|
||||
import httpx
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
||||
from nanobot.providers.oauth_utils import get_auth_headers
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse, LongContextError, ToolCallRequest
|
||||
from nanobot.providers.oauth_utils import get_auth_headers, get_claude_code_system_prefix
|
||||
|
||||
|
||||
class AnthropicOAuthProvider(LLMProvider):
|
||||
@@ -27,7 +27,7 @@ class AnthropicOAuthProvider(LLMProvider):
|
||||
def __init__(
|
||||
self,
|
||||
oauth_token: str,
|
||||
default_model: str = "claude-opus-4-5",
|
||||
default_model: str = "claude-opus-4-7",
|
||||
api_base: str | None = None,
|
||||
thinking_budget: int = 0,
|
||||
):
|
||||
@@ -51,17 +51,91 @@ class AnthropicOAuthProvider(LLMProvider):
|
||||
def _normalize_model(model: str) -> str:
|
||||
"""Normalize model name for the Anthropic API.
|
||||
|
||||
Anthropic model IDs use hyphens (claude-sonnet-4-5), but users often
|
||||
write dots (claude-sonnet-4.5). Normalize so both work.
|
||||
Anthropic model IDs use hyphens (claude-sonnet-4-6), but users often
|
||||
write dots (claude-sonnet-4.6). Normalize so both work.
|
||||
"""
|
||||
return model.replace(".", "-")
|
||||
|
||||
async def _get_client(self) -> httpx.AsyncClient:
|
||||
"""Get or create async HTTP client."""
|
||||
if self._client is None:
|
||||
self._client = httpx.AsyncClient(timeout=300.0)
|
||||
self._client = httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(300.0, pool=30.0),
|
||||
)
|
||||
return self._client
|
||||
|
||||
async def _reset_client(self) -> None:
|
||||
"""Destroy and recreate the HTTP client after connection errors."""
|
||||
old = self._client
|
||||
self._client = None
|
||||
if old:
|
||||
try:
|
||||
await old.aclose()
|
||||
except Exception:
|
||||
pass
|
||||
logger.warning("Reset httpx client (pool recycled)")
|
||||
|
||||
async def _diagnose_connectivity(self) -> None:
|
||||
"""Run diagnostics when ConnectTimeout occurs to understand why."""
|
||||
import socket
|
||||
import asyncio
|
||||
|
||||
# 1. Raw socket test (bypasses httpx entirely)
|
||||
try:
|
||||
t0 = __import__('time').monotonic()
|
||||
s = socket.create_connection(('api.anthropic.com', 443), timeout=10)
|
||||
elapsed = __import__('time').monotonic() - t0
|
||||
s.close()
|
||||
logger.warning(f"DIAG: raw socket connect OK in {elapsed:.3f}s")
|
||||
except Exception as e:
|
||||
logger.error(f"DIAG: raw socket connect FAILED: {e}")
|
||||
|
||||
# 2. asyncio connect test (same event loop)
|
||||
try:
|
||||
t0 = __import__('time').monotonic()
|
||||
reader, writer = await asyncio.wait_for(
|
||||
asyncio.open_connection('api.anthropic.com', 443),
|
||||
timeout=10.0,
|
||||
)
|
||||
elapsed = __import__('time').monotonic() - t0
|
||||
writer.close()
|
||||
await writer.wait_closed()
|
||||
logger.warning(f"DIAG: asyncio connect OK in {elapsed:.3f}s")
|
||||
except Exception as e:
|
||||
logger.error(f"DIAG: asyncio connect FAILED: {e}")
|
||||
|
||||
# 3. Fresh httpx client test (new pool)
|
||||
try:
|
||||
t0 = __import__('time').monotonic()
|
||||
async with httpx.AsyncClient(timeout=10.0) as fresh:
|
||||
r = await fresh.get('https://api.anthropic.com/')
|
||||
elapsed = __import__('time').monotonic() - t0
|
||||
logger.warning(f"DIAG: fresh httpx OK in {elapsed:.3f}s (status={r.status_code})")
|
||||
except Exception as e:
|
||||
logger.error(f"DIAG: fresh httpx FAILED: {e}")
|
||||
|
||||
# 4. DNS resolution
|
||||
try:
|
||||
ips = socket.getaddrinfo('api.anthropic.com', 443)
|
||||
logger.warning(f"DIAG: DNS resolved to {len(ips)} entries, first={ips[0][4][0]}")
|
||||
except Exception as e:
|
||||
logger.error(f"DIAG: DNS FAILED: {e}")
|
||||
|
||||
# 5. Connection pool state of the broken client
|
||||
if self._client:
|
||||
transport = self._client._transport
|
||||
if hasattr(transport, '_pool'):
|
||||
pool = transport._pool
|
||||
conns = getattr(pool, '_connections', [])
|
||||
reqs = getattr(pool, '_requests', [])
|
||||
logger.warning(
|
||||
f"DIAG: pool state: {len(conns)} connections, "
|
||||
f"{len(reqs)} pending requests"
|
||||
)
|
||||
for i, conn in enumerate(conns[:5]):
|
||||
state = getattr(conn, '_state', 'unknown')
|
||||
logger.warning(f"DIAG: conn[{i}] state={state}")
|
||||
|
||||
def _prepare_messages(
|
||||
self,
|
||||
messages: list[dict[str, Any]]
|
||||
@@ -252,18 +326,35 @@ class AnthropicOAuthProvider(LLMProvider):
|
||||
"""Make request to Anthropic API."""
|
||||
client = await self._get_client()
|
||||
|
||||
# Cache the last user message so conversation history is cached across turns
|
||||
if messages:
|
||||
last = messages[-1]
|
||||
if last.get("role") == "user":
|
||||
content = last["content"]
|
||||
if isinstance(content, str):
|
||||
last = {**last, "content": [{"type": "text", "text": content, "cache_control": {"type": "ephemeral"}}]}
|
||||
elif isinstance(content, list) and content:
|
||||
new_content = list(content)
|
||||
new_content[-1] = {**new_content[-1], "cache_control": {"type": "ephemeral"}}
|
||||
last = {**last, "content": new_content}
|
||||
messages = messages[:-1] + [last]
|
||||
# Add cache breakpoints on the last TWO user messages (4-breakpoint strategy):
|
||||
# BP3: Second-to-last user message (stable history from previous turn)
|
||||
# BP4: Last user message (current turn, will become BP3 next turn)
|
||||
# This allows BP3 to reuse what BP4 cached last turn.
|
||||
user_indices = [i for i, m in enumerate(messages) if m.get("role") == "user"]
|
||||
|
||||
if len(user_indices) >= 2:
|
||||
# BP3: Second-to-last user message
|
||||
idx = user_indices[-2]
|
||||
msg = messages[idx]
|
||||
content = msg["content"]
|
||||
if isinstance(content, str):
|
||||
messages[idx] = {**msg, "content": [{"type": "text", "text": content, "cache_control": {"type": "ephemeral"}}]}
|
||||
elif isinstance(content, list) and content:
|
||||
new_content = list(content)
|
||||
new_content[-1] = {**new_content[-1], "cache_control": {"type": "ephemeral"}}
|
||||
messages[idx] = {**msg, "content": new_content}
|
||||
|
||||
if len(user_indices) >= 1:
|
||||
# BP4: Last user message
|
||||
idx = user_indices[-1]
|
||||
msg = messages[idx]
|
||||
content = msg["content"]
|
||||
if isinstance(content, str):
|
||||
messages[idx] = {**msg, "content": [{"type": "text", "text": content, "cache_control": {"type": "ephemeral"}}]}
|
||||
elif isinstance(content, list) and content:
|
||||
new_content = list(content)
|
||||
new_content[-1] = {**new_content[-1], "cache_control": {"type": "ephemeral"}}
|
||||
messages[idx] = {**msg, "content": new_content}
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
@@ -286,7 +377,14 @@ class AnthropicOAuthProvider(LLMProvider):
|
||||
payload["temperature"] = temperature
|
||||
|
||||
if system:
|
||||
payload["system"] = [{"type": "text", "text": system, "cache_control": {"type": "ephemeral", "ttl": "1h"}}]
|
||||
payload["system"] = [
|
||||
{"type": "text", "text": get_claude_code_system_prefix()},
|
||||
{"type": "text", "text": system, "cache_control": {"type": "ephemeral", "ttl": "1h"}},
|
||||
]
|
||||
else:
|
||||
payload["system"] = [
|
||||
{"type": "text", "text": get_claude_code_system_prefix()},
|
||||
]
|
||||
|
||||
if tools:
|
||||
cached_tools = list(tools)
|
||||
@@ -316,50 +414,131 @@ class AnthropicOAuthProvider(LLMProvider):
|
||||
headers.get("anthropic-beta", "none"),
|
||||
)
|
||||
|
||||
response = await client.post(
|
||||
self._get_api_url(),
|
||||
headers=headers,
|
||||
json=payload,
|
||||
)
|
||||
# Debug: Log tool names for diagnostic purposes
|
||||
if payload.get("tools"):
|
||||
tool_names = [t.get("name", "unnamed") for t in payload["tools"]]
|
||||
logger.debug(f"Tool names in request: {tool_names}")
|
||||
|
||||
# Dump rate limit headers for analysis
|
||||
try:
|
||||
import datetime
|
||||
import os
|
||||
header_dump = {
|
||||
"timestamp": datetime.datetime.utcnow().isoformat(),
|
||||
"status_code": response.status_code,
|
||||
"model": payload.get("model"),
|
||||
"headers": dict(response.headers),
|
||||
}
|
||||
dump_path = "/root/.nanobot/workspace/api_headers.jsonl"
|
||||
with open(dump_path, "a") as f:
|
||||
f.write(json.dumps(header_dump) + "\n")
|
||||
# Capture rate limit state for quota-based model switching
|
||||
hdrs = response.headers
|
||||
rate_limit_state = {
|
||||
"updated_at": datetime.datetime.utcnow().isoformat(),
|
||||
"model": payload.get("model"),
|
||||
"weekly_all_models": float(hdrs["anthropic-ratelimit-unified-7d-utilization"]) if hdrs.get("anthropic-ratelimit-unified-7d-utilization") else None,
|
||||
"weekly_sonnet": float(hdrs["anthropic-ratelimit-unified-7d_sonnet-utilization"]) if hdrs.get("anthropic-ratelimit-unified-7d_sonnet-utilization") else None,
|
||||
"session_5h": float(hdrs["anthropic-ratelimit-unified-5h-utilization"]) if hdrs.get("anthropic-ratelimit-unified-5h-utilization") else None,
|
||||
"weekly_reset": int(hdrs["anthropic-ratelimit-unified-7d-reset"]) if hdrs.get("anthropic-ratelimit-unified-7d-reset") else None,
|
||||
"session_reset": int(hdrs["anthropic-ratelimit-unified-5h-reset"]) if hdrs.get("anthropic-ratelimit-unified-5h-reset") else None,
|
||||
"binding_limit": hdrs.get("anthropic-ratelimit-unified-representative-claim"),
|
||||
"sonnet_fallback": hdrs.get("anthropic-ratelimit-unified-fallback"),
|
||||
}
|
||||
state_path = "/root/.nanobot/workspace/memory/rate_limits.json"
|
||||
os.makedirs(os.path.dirname(state_path), exist_ok=True)
|
||||
with open(state_path, "w") as f:
|
||||
json.dump(rate_limit_state, f, indent=2)
|
||||
except Exception as e:
|
||||
logger.warning("Rate limit header capture failed: {}", e)
|
||||
# Debug: Log message structure to diagnose orphaned tool_result errors
|
||||
for idx, m in enumerate(payload.get("messages", [])):
|
||||
role = m.get("role", "?")
|
||||
content = m.get("content", "")
|
||||
if isinstance(content, list):
|
||||
block_types = [b.get("type", "?") for b in content]
|
||||
logger.debug(f" msg[{idx}] role={role} blocks={block_types}")
|
||||
else:
|
||||
logger.debug(f" msg[{idx}] role={role} text={str(content)[:80]}")
|
||||
|
||||
if response.status_code != 200:
|
||||
error_text = response.text
|
||||
raise Exception(f"Anthropic API error {response.status_code}: {error_text}")
|
||||
import asyncio
|
||||
import time as _time
|
||||
|
||||
return response.json()
|
||||
max_retries = 3
|
||||
base_delay = 2.0 # seconds
|
||||
|
||||
for attempt in range(max_retries + 1):
|
||||
_t0 = _time.monotonic()
|
||||
try:
|
||||
response = await client.post(
|
||||
self._get_api_url(),
|
||||
headers=headers,
|
||||
json=payload,
|
||||
)
|
||||
except httpx.ConnectTimeout:
|
||||
elapsed = _time.monotonic() - _t0
|
||||
logger.error(f"ConnectTimeout after {elapsed:.1f}s (attempt {attempt+1}/{max_retries+1})")
|
||||
if attempt == 0:
|
||||
await self._diagnose_connectivity()
|
||||
await self._reset_client()
|
||||
if attempt < max_retries:
|
||||
delay = base_delay * (2 ** attempt)
|
||||
logger.info(f"Retrying in {delay:.1f}s...")
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
raise
|
||||
except httpx.PoolTimeout:
|
||||
elapsed = _time.monotonic() - _t0
|
||||
logger.error(f"PoolTimeout after {elapsed:.1f}s (attempt {attempt+1}/{max_retries+1})")
|
||||
await self._reset_client()
|
||||
if attempt < max_retries:
|
||||
delay = base_delay * (2 ** attempt)
|
||||
logger.info(f"Retrying in {delay:.1f}s...")
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
raise
|
||||
except (httpx.ConnectError, httpx.TimeoutException) as e:
|
||||
elapsed = _time.monotonic() - _t0
|
||||
logger.error(f"{type(e).__name__} after {elapsed:.1f}s (attempt {attempt+1}/{max_retries+1})")
|
||||
if attempt < max_retries:
|
||||
delay = base_delay * (2 ** attempt)
|
||||
logger.info(f"Retrying in {delay:.1f}s...")
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
raise
|
||||
elapsed = _time.monotonic() - _t0
|
||||
if elapsed > 30:
|
||||
logger.warning(f"Anthropic API slow response: {elapsed:.1f}s")
|
||||
|
||||
# Dump rate limit headers for analysis
|
||||
try:
|
||||
import datetime
|
||||
import os
|
||||
header_dump = {
|
||||
"timestamp": datetime.datetime.now(datetime.UTC).isoformat(),
|
||||
"status_code": response.status_code,
|
||||
"model": payload.get("model"),
|
||||
"headers": dict(response.headers),
|
||||
}
|
||||
dump_path = "/root/.nanobot/workspace/api_headers.jsonl"
|
||||
with open(dump_path, "a") as f:
|
||||
f.write(json.dumps(header_dump) + "\n")
|
||||
# Capture rate limit state for quota-based model switching
|
||||
hdrs = response.headers
|
||||
rate_limit_state = {
|
||||
"updated_at": datetime.datetime.utcnow().isoformat(),
|
||||
"model": payload.get("model"),
|
||||
"weekly_all_models": float(hdrs["anthropic-ratelimit-unified-7d-utilization"]) if hdrs.get("anthropic-ratelimit-unified-7d-utilization") else None,
|
||||
"weekly_sonnet": float(hdrs["anthropic-ratelimit-unified-7d_sonnet-utilization"]) if hdrs.get("anthropic-ratelimit-unified-7d_sonnet-utilization") else None,
|
||||
"session_5h": float(hdrs["anthropic-ratelimit-unified-5h-utilization"]) if hdrs.get("anthropic-ratelimit-unified-5h-utilization") else None,
|
||||
"weekly_reset": int(hdrs["anthropic-ratelimit-unified-7d-reset"]) if hdrs.get("anthropic-ratelimit-unified-7d-reset") else None,
|
||||
"session_reset": int(hdrs["anthropic-ratelimit-unified-5h-reset"]) if hdrs.get("anthropic-ratelimit-unified-5h-reset") else None,
|
||||
"binding_limit": hdrs.get("anthropic-ratelimit-unified-representative-claim"),
|
||||
"sonnet_fallback": hdrs.get("anthropic-ratelimit-unified-fallback"),
|
||||
}
|
||||
state_path = "/root/.nanobot/workspace/memory/rate_limits.json"
|
||||
os.makedirs(os.path.dirname(state_path), exist_ok=True)
|
||||
with open(state_path, "w") as f:
|
||||
json.dump(rate_limit_state, f, indent=2)
|
||||
except Exception as e:
|
||||
logger.warning("Rate limit header capture failed: {}", e)
|
||||
|
||||
# Retry on 5xx server errors and 429 rate limits
|
||||
if response.status_code >= 500 or response.status_code == 429:
|
||||
error_text = response.text
|
||||
logger.warning(f"Anthropic API {response.status_code} (attempt {attempt+1}/{max_retries+1}): {error_text[:200]}")
|
||||
|
||||
# Long context 429 — retrying won't help, need to trim context
|
||||
if response.status_code == 429 and "long context" in error_text.lower():
|
||||
raise LongContextError(f"Context too long for current plan: {error_text[:200]}")
|
||||
|
||||
if attempt < max_retries:
|
||||
if response.status_code == 429:
|
||||
retry_after = response.headers.get("retry-after")
|
||||
delay = float(retry_after) if retry_after else base_delay * (2 ** attempt)
|
||||
else:
|
||||
delay = base_delay * (2 ** attempt)
|
||||
logger.info(f"Retrying in {delay:.1f}s...")
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
raise Exception(f"Anthropic API error {response.status_code}: {error_text}")
|
||||
|
||||
if response.status_code != 200:
|
||||
error_text = response.text
|
||||
raise Exception(f"Anthropic API error {response.status_code}: {error_text}")
|
||||
|
||||
return response.json()
|
||||
|
||||
# Should not reach here, but just in case
|
||||
raise Exception("Exhausted all retry attempts")
|
||||
|
||||
async def chat(
|
||||
self,
|
||||
@@ -378,7 +557,7 @@ class AnthropicOAuthProvider(LLMProvider):
|
||||
if "/" in model:
|
||||
model = model.split("/")[-1]
|
||||
|
||||
# Normalize dots to hyphens (claude-sonnet-4.5 -> claude-sonnet-4-5)
|
||||
# Normalize dots to hyphens (claude-sonnet-4.6 -> claude-sonnet-4-6)
|
||||
model = self._normalize_model(model)
|
||||
|
||||
system, prepared_messages = self._prepare_messages(messages)
|
||||
@@ -411,9 +590,13 @@ class AnthropicOAuthProvider(LLMProvider):
|
||||
beta_flags=beta_flags,
|
||||
)
|
||||
return self._parse_response(response)
|
||||
except LongContextError:
|
||||
raise # Let caller handle context trimming
|
||||
except Exception as e:
|
||||
logger.exception("Exception in chat():")
|
||||
error_msg = f"{type(e).__name__}: {str(e)}" if str(e) else f"{type(e).__name__} (no message)"
|
||||
return LLMResponse(
|
||||
content=f"Error calling LLM: {str(e)}",
|
||||
content=f"Error calling LLM: {error_msg}",
|
||||
finish_reason="error",
|
||||
)
|
||||
|
||||
|
||||
@@ -28,6 +28,11 @@ class LLMResponse:
|
||||
return len(self.tool_calls) > 0
|
||||
|
||||
|
||||
class LongContextError(Exception):
|
||||
"""Raised when the API rejects a request due to long context limits."""
|
||||
pass
|
||||
|
||||
|
||||
class LLMProvider(ABC):
|
||||
"""
|
||||
Abstract base class for LLM providers.
|
||||
|
||||
@@ -37,7 +37,7 @@ class LiteLLMProvider(LLMProvider):
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
default_model: str = "anthropic/claude-opus-4-5",
|
||||
default_model: str = "anthropic/claude-opus-4-7",
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
provider_name: str | None = None,
|
||||
):
|
||||
@@ -187,7 +187,7 @@ class LiteLLMProvider(LLMProvider):
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
tools: Optional list of tool definitions in OpenAI format.
|
||||
model: Model identifier (e.g., 'anthropic/claude-sonnet-4-5').
|
||||
model: Model identifier (e.g., 'anthropic/claude-sonnet-4-6').
|
||||
max_tokens: Maximum tokens in response.
|
||||
temperature: Sampling temperature.
|
||||
|
||||
|
||||
@@ -36,3 +36,11 @@ def get_auth_headers(token: str, is_oauth: bool = False) -> dict[str, str]:
|
||||
headers["x-api-key"] = token
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
def get_claude_code_system_prefix() -> str:
|
||||
"""Get the required system prompt prefix for OAuth tokens.
|
||||
|
||||
Anthropic requires this identity declaration for OAuth auth.
|
||||
"""
|
||||
return "You are a Claude agent, built on Anthropic's Claude Agent SDK."
|
||||
|
||||
+77
-50
@@ -1,6 +1,7 @@
|
||||
"""Session management for conversation history."""
|
||||
|
||||
import json
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
@@ -15,10 +16,12 @@ from nanobot.utils.helpers import ensure_dir, safe_filename
|
||||
class Session:
|
||||
"""
|
||||
A conversation session.
|
||||
|
||||
|
||||
Stores messages in JSONL format for easy reading and persistence.
|
||||
|
||||
Messages are trimmed after consolidation to keep session size manageable.
|
||||
"""
|
||||
|
||||
|
||||
key: str # channel:chat_id
|
||||
messages: list[dict[str, Any]] = field(default_factory=list)
|
||||
created_at: datetime = field(default_factory=datetime.now)
|
||||
@@ -56,16 +59,25 @@ class Session:
|
||||
trimming old tool chains safely at token thresholds, so we send the full
|
||||
history and let the server decide what to drop.
|
||||
|
||||
Messages with ``_hidden_sig`` get a ``[HIDDEN:{sig}]`` prefix applied to
|
||||
their content so the model knows the user never saw them. The prefix is
|
||||
applied at read time (not stored in content) to preserve prompt-cache
|
||||
stability: the same prefixed string is produced every turn.
|
||||
|
||||
Returns:
|
||||
List of messages in LLM format (API-relevant fields only).
|
||||
"""
|
||||
return [
|
||||
{k: v for k, v in m.items() if k in self._API_FIELDS and v is not None}
|
||||
for m in self.messages
|
||||
]
|
||||
out: list[dict[str, Any]] = []
|
||||
for m in self.messages:
|
||||
msg = {k: v for k, v in m.items() if k in self._API_FIELDS and v is not None}
|
||||
sig = m.get("_hidden_sig")
|
||||
if sig and isinstance(msg.get("content"), str):
|
||||
msg["content"] = f"[HIDDEN:{sig}] {msg['content']}"
|
||||
out.append(msg)
|
||||
return out
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Clear all messages in the session."""
|
||||
"""Clear all messages and reset session to initial state."""
|
||||
self.messages = []
|
||||
self.updated_at = datetime.now()
|
||||
|
||||
@@ -73,19 +85,25 @@ class Session:
|
||||
class SessionManager:
|
||||
"""
|
||||
Manages conversation sessions.
|
||||
|
||||
|
||||
Sessions are stored as JSONL files in the sessions directory.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self, workspace: Path):
|
||||
self.workspace = workspace
|
||||
self.sessions_dir = ensure_dir(Path.home() / ".nanobot" / "sessions")
|
||||
self.sessions_dir = ensure_dir(self.workspace / "sessions")
|
||||
self.legacy_sessions_dir = Path.home() / ".nanobot" / "sessions"
|
||||
self._cache: dict[str, Session] = {}
|
||||
|
||||
def _get_session_path(self, key: str) -> Path:
|
||||
"""Get the file path for a session."""
|
||||
safe_key = safe_filename(key.replace(":", "_"))
|
||||
return self.sessions_dir / f"{safe_key}.jsonl"
|
||||
|
||||
def _get_legacy_session_path(self, key: str) -> Path:
|
||||
"""Legacy global session path (~/.nanobot/sessions/)."""
|
||||
safe_key = safe_filename(key.replace(":", "_"))
|
||||
return self.legacy_sessions_dir / f"{safe_key}.jsonl"
|
||||
|
||||
def get_or_create(self, key: str) -> Session:
|
||||
"""
|
||||
@@ -97,11 +115,9 @@ class SessionManager:
|
||||
Returns:
|
||||
The session.
|
||||
"""
|
||||
# Check cache
|
||||
if key in self._cache:
|
||||
return self._cache[key]
|
||||
|
||||
# Try to load from disk
|
||||
session = self._load(key)
|
||||
if session is None:
|
||||
session = Session(key=key)
|
||||
@@ -112,78 +128,88 @@ class SessionManager:
|
||||
def _load(self, key: str) -> Session | None:
|
||||
"""Load a session from disk."""
|
||||
path = self._get_session_path(key)
|
||||
|
||||
if not path.exists():
|
||||
legacy_path = self._get_legacy_session_path(key)
|
||||
if legacy_path.exists():
|
||||
try:
|
||||
shutil.move(str(legacy_path), str(path))
|
||||
logger.info("Migrated session {} from legacy path", key)
|
||||
except Exception:
|
||||
logger.exception("Failed to migrate session {}", key)
|
||||
|
||||
if not path.exists():
|
||||
return None
|
||||
|
||||
|
||||
try:
|
||||
messages = []
|
||||
metadata = {}
|
||||
created_at = None
|
||||
|
||||
with open(path) as f:
|
||||
|
||||
with open(path, encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
|
||||
|
||||
data = json.loads(line)
|
||||
|
||||
|
||||
if data.get("_type") == "metadata":
|
||||
metadata = data.get("metadata", {})
|
||||
created_at = datetime.fromisoformat(data["created_at"]) if data.get("created_at") else None
|
||||
# Ignore legacy last_consolidated field
|
||||
else:
|
||||
messages.append(data)
|
||||
|
||||
|
||||
return Session(
|
||||
key=key,
|
||||
messages=messages,
|
||||
created_at=created_at or datetime.now(),
|
||||
metadata=metadata
|
||||
metadata=metadata,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load session {key}: {e}")
|
||||
logger.warning("Failed to load session {}: {}", key, e)
|
||||
return None
|
||||
|
||||
def save(self, session: Session) -> None:
|
||||
"""Save a session to disk."""
|
||||
path = self._get_session_path(session.key)
|
||||
|
||||
with open(path, "w") as f:
|
||||
# Write metadata first
|
||||
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
metadata_line = {
|
||||
"_type": "metadata",
|
||||
"key": session.key,
|
||||
"created_at": session.created_at.isoformat(),
|
||||
"updated_at": session.updated_at.isoformat(),
|
||||
"metadata": session.metadata
|
||||
"metadata": session.metadata,
|
||||
}
|
||||
f.write(json.dumps(metadata_line) + "\n")
|
||||
|
||||
# Write messages
|
||||
f.write(json.dumps(metadata_line, ensure_ascii=False) + "\n")
|
||||
for msg in session.messages:
|
||||
f.write(json.dumps(msg) + "\n")
|
||||
|
||||
f.write(json.dumps(msg, ensure_ascii=False) + "\n")
|
||||
|
||||
self._cache[session.key] = session
|
||||
self._append_audit(session)
|
||||
|
||||
def _append_audit(self, session: Session) -> None:
|
||||
"""Append session state to an audit log (append-only, rotated monthly)."""
|
||||
now = datetime.now()
|
||||
safe_key = safe_filename(session.key.replace(":", "_"))
|
||||
audit_path = self.sessions_dir / f"{safe_key}.audit.{now:%Y-%m}.jsonl"
|
||||
try:
|
||||
with open(audit_path, "a", encoding="utf-8") as f:
|
||||
marker = {
|
||||
"_type": "save_marker",
|
||||
"timestamp": now.isoformat(),
|
||||
"message_count": len(session.messages),
|
||||
}
|
||||
f.write(json.dumps(marker, ensure_ascii=False) + "\n")
|
||||
for msg in session.messages:
|
||||
f.write(json.dumps(msg, ensure_ascii=False) + "\n")
|
||||
except Exception as e:
|
||||
logger.warning("Audit log write failed for {}: {}", session.key, e)
|
||||
|
||||
def delete(self, key: str) -> bool:
|
||||
"""
|
||||
Delete a session.
|
||||
|
||||
Args:
|
||||
key: Session key.
|
||||
|
||||
Returns:
|
||||
True if deleted, False if not found.
|
||||
"""
|
||||
# Remove from cache
|
||||
def invalidate(self, key: str) -> None:
|
||||
"""Remove a session from the in-memory cache."""
|
||||
self._cache.pop(key, None)
|
||||
|
||||
# Remove file
|
||||
path = self._get_session_path(key)
|
||||
if path.exists():
|
||||
path.unlink()
|
||||
return True
|
||||
return False
|
||||
|
||||
def list_sessions(self) -> list[dict[str, Any]]:
|
||||
"""
|
||||
@@ -197,13 +223,14 @@ class SessionManager:
|
||||
for path in self.sessions_dir.glob("*.jsonl"):
|
||||
try:
|
||||
# Read just the metadata line
|
||||
with open(path) as f:
|
||||
with open(path, encoding="utf-8") as f:
|
||||
first_line = f.readline().strip()
|
||||
if first_line:
|
||||
data = json.loads(first_line)
|
||||
if data.get("_type") == "metadata":
|
||||
key = data.get("key") or path.stem.replace("_", ":", 1)
|
||||
sessions.append({
|
||||
"key": path.stem.replace("_", ":"),
|
||||
"key": key,
|
||||
"created_at": data.get("created_at"),
|
||||
"updated_at": data.get("updated_at"),
|
||||
"path": str(path)
|
||||
|
||||
@@ -47,6 +47,14 @@ dev = [
|
||||
"pytest-asyncio>=0.21.0",
|
||||
"ruff>=0.1.0",
|
||||
]
|
||||
mem0 = [
|
||||
"mem0ai>=0.1.0",
|
||||
]
|
||||
matrix = [
|
||||
"matrix-nio>=0.20.0",
|
||||
"mistune>=3.0.0",
|
||||
"nh3>=0.2.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
nanobot = "nanobot.cli.commands:app"
|
||||
|
||||
Executable
+34
@@ -0,0 +1,34 @@
|
||||
#!/bin/bash
|
||||
# test-pr.sh - Quick PR testing script for nanobot staging
|
||||
#
|
||||
# Usage: ./test-pr.sh <pr-number> [test-message]
|
||||
# Example: ./test-pr.sh 31 "test tool use feature"
|
||||
|
||||
set -e
|
||||
|
||||
PR_NUM="$1"
|
||||
TEST_MSG="${2:-Hello, testing PR #$PR_NUM}"
|
||||
REPO_DIR="/config/workspace/nanobot-oauth-port/nanobot-fork"
|
||||
STAGING_CONFIG="/config/workspace/.nanobot-staging/config.json"
|
||||
|
||||
if [ -z "$PR_NUM" ]; then
|
||||
echo "Usage: $0 <pr-number> [test-message]"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "==> Fetching PR #$PR_NUM..."
|
||||
cd "$REPO_DIR"
|
||||
git fetch wylab "+pull/$PR_NUM/head:pr-$PR_NUM"
|
||||
|
||||
echo "==> Checking out pr-$PR_NUM..."
|
||||
git checkout "pr-$PR_NUM"
|
||||
|
||||
echo "==> Installing in editable mode..."
|
||||
uv pip install -e . -q
|
||||
|
||||
echo "==> Testing with message: $TEST_MSG"
|
||||
NANOBOT_CONFIG="$STAGING_CONFIG" "$REPO_DIR/.venv/bin/nanobot" agent -m "$TEST_MSG"
|
||||
|
||||
echo ""
|
||||
echo "==> Test complete. Branch pr-$PR_NUM is still checked out."
|
||||
echo " Run 'git checkout main' to return to main branch."
|
||||
@@ -28,7 +28,7 @@ def mock_session_manager():
|
||||
"messages": [],
|
||||
"metadata": {},
|
||||
})
|
||||
session_mgr.save = AsyncMock()
|
||||
session_mgr.save = MagicMock() # Synchronous in production, not async
|
||||
return session_mgr
|
||||
|
||||
|
||||
|
||||
@@ -10,14 +10,14 @@ def provider():
|
||||
"""Create provider with test OAuth token."""
|
||||
return AnthropicOAuthProvider(
|
||||
oauth_token="sk-ant-oat01-test-token",
|
||||
default_model="claude-opus-4-5"
|
||||
default_model="claude-opus-4-7"
|
||||
)
|
||||
|
||||
|
||||
def test_provider_init(provider):
|
||||
"""Provider should initialize with OAuth token."""
|
||||
assert provider.oauth_token == "sk-ant-oat01-test-token"
|
||||
assert provider.default_model == "claude-opus-4-5"
|
||||
assert provider.default_model == "claude-opus-4-7"
|
||||
|
||||
|
||||
def test_provider_uses_bearer_auth(provider):
|
||||
@@ -28,18 +28,8 @@ def test_provider_uses_bearer_auth(provider):
|
||||
assert "x-api-key" not in headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_prepends_system_prompt(provider):
|
||||
"""Chat should prepend Claude Code identity to system prompt."""
|
||||
messages = [{"role": "user", "content": "Hello"}]
|
||||
|
||||
with patch.object(provider, "_make_request", new_callable=AsyncMock) as mock:
|
||||
mock.return_value = {"content": [{"type": "text", "text": "Hi"}], "stop_reason": "end_turn"}
|
||||
await provider.chat(messages)
|
||||
|
||||
call_args = mock.call_args
|
||||
system = call_args[1]["system"]
|
||||
assert "Claude Code" in system
|
||||
# test_chat_prepends_system_prompt removed - feature no longer exists
|
||||
# System prompt handling is done by the agent loop, not the provider
|
||||
|
||||
|
||||
def test_parse_response_text(provider):
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Test that bash tool handles heredoc commands correctly.
|
||||
|
||||
Reproduces the bug where `; echo '<<exit>>'` appended on the same line
|
||||
as a heredoc terminator prevents bash from recognizing the terminator,
|
||||
causing the session to hang forever.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import pytest
|
||||
from nanobot.agent.tools.anthropic.bash import BashTool20250124
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_heredoc_command():
|
||||
"""Heredoc commands must complete without hanging."""
|
||||
tool = BashTool20250124()
|
||||
|
||||
# Simple command works
|
||||
result = await tool(command="echo hello")
|
||||
assert result.output == "hello"
|
||||
|
||||
# Heredoc command — this is the exact pattern that caused the hang
|
||||
result = await asyncio.wait_for(
|
||||
tool(command="cat << 'EOF'\nline1\nline2\nEOF"),
|
||||
timeout=5.0,
|
||||
)
|
||||
assert "line1" in result.output
|
||||
assert "line2" in result.output
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_heredoc_append_to_file():
|
||||
"""Heredoc append (the exact pattern the LLM uses) must work."""
|
||||
tool = BashTool20250124()
|
||||
|
||||
result = await asyncio.wait_for(
|
||||
tool(command="cat >> /tmp/test_heredoc_bash.txt << 'EOF'\nhello world\nEOF"),
|
||||
timeout=5.0,
|
||||
)
|
||||
# Should complete without error
|
||||
assert result.error is None or result.error == ""
|
||||
|
||||
# Verify the file was written
|
||||
result2 = await tool(command="cat /tmp/test_heredoc_bash.txt")
|
||||
assert "hello world" in result2.output
|
||||
|
||||
# Cleanup
|
||||
await tool(command="rm -f /tmp/test_heredoc_bash.txt")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_regular_commands_still_work():
|
||||
"""Ensure regular commands still work after the fix."""
|
||||
tool = BashTool20250124()
|
||||
|
||||
# Semicolons in commands
|
||||
result = await tool(command="echo a; echo b")
|
||||
assert "a" in result.output
|
||||
assert "b" in result.output
|
||||
|
||||
# Multiline script
|
||||
result = await tool(command="for i in 1 2 3; do echo $i; done")
|
||||
assert "1" in result.output
|
||||
assert "3" in result.output
|
||||
|
||||
# Command with exit code
|
||||
result = await tool(command="true")
|
||||
assert result.output == "(no output)" or result.output is not None
|
||||
@@ -40,7 +40,7 @@ async def test_bash_tool_restart():
|
||||
|
||||
# Restart
|
||||
result = await tool(restart=True)
|
||||
assert "restarted" in result.output.lower()
|
||||
assert "restarted" in (result.system or result.output or "").lower()
|
||||
|
||||
# Variable should be gone
|
||||
result2 = await tool(command="echo $TEST_VAR")
|
||||
|
||||
@@ -50,11 +50,12 @@ async def test_beta_flags_collected_from_tools():
|
||||
tools=tools_with_flags
|
||||
)
|
||||
|
||||
# Check that beta flag was added to headers
|
||||
# Check that beta flag was added to headers (merged with hardcoded flags)
|
||||
call_args = mock_client.post.call_args
|
||||
headers = call_args[1]["headers"]
|
||||
assert "anthropic-beta" in headers
|
||||
assert headers["anthropic-beta"] == "computer-use-2025-11-24"
|
||||
# Should include hardcoded flags + tool flag, sorted alphabetically
|
||||
assert headers["anthropic-beta"] == "claude-code-20250219,computer-use-2025-11-24,context-management-2025-06-27,oauth-2025-04-20"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -99,5 +100,5 @@ async def test_multiple_beta_flags_joined():
|
||||
call_args = mock_client.post.call_args
|
||||
headers = call_args[1]["headers"]
|
||||
assert "anthropic-beta" in headers
|
||||
# Should be sorted alphabetically and joined with comma
|
||||
assert headers["anthropic-beta"] == "flag-a,flag-b"
|
||||
# Should include hardcoded flags + tool flags, sorted alphabetically and joined with comma
|
||||
assert headers["anthropic-beta"] == "claude-code-20250219,context-management-2025-06-27,flag-a,flag-b,oauth-2025-04-20"
|
||||
|
||||
+11
-11
@@ -29,6 +29,7 @@ def mock_paths():
|
||||
|
||||
config_file = base_dir / "config.json"
|
||||
workspace_dir = base_dir / "workspace"
|
||||
workspace_dir.mkdir() # Create workspace directory
|
||||
|
||||
mock_cp.return_value = config_file
|
||||
mock_ws.return_value = workspace_dir
|
||||
@@ -56,21 +57,20 @@ def test_onboard_fresh_install(mock_paths):
|
||||
|
||||
|
||||
def test_onboard_existing_config_refresh(mock_paths):
|
||||
"""Config exists, user declines overwrite — should refresh (load-merge-save)."""
|
||||
"""Config exists, user declines overwrite — should exit without changes."""
|
||||
config_file, workspace_dir = mock_paths
|
||||
config_file.write_text('{"existing": true}')
|
||||
|
||||
result = runner.invoke(app, ["onboard"], input="n\n")
|
||||
|
||||
# User declined, so command exits (typer.Exit() returns 0)
|
||||
assert result.exit_code == 0
|
||||
assert "Config already exists" in result.stdout
|
||||
assert "existing values preserved" in result.stdout
|
||||
assert workspace_dir.exists()
|
||||
assert (workspace_dir / "AGENTS.md").exists()
|
||||
assert "Overwrite?" in result.stdout
|
||||
|
||||
|
||||
def test_onboard_existing_config_overwrite(mock_paths):
|
||||
"""Config exists, user confirms overwrite — should reset to defaults."""
|
||||
"""Config exists, user confirms overwrite — should create new config."""
|
||||
config_file, workspace_dir = mock_paths
|
||||
config_file.write_text('{"existing": true}')
|
||||
|
||||
@@ -78,20 +78,20 @@ def test_onboard_existing_config_overwrite(mock_paths):
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert "Config already exists" in result.stdout
|
||||
assert "Config reset to defaults" in result.stdout
|
||||
assert "Created config" in result.stdout
|
||||
assert workspace_dir.exists()
|
||||
|
||||
|
||||
def test_onboard_existing_workspace_safe_create(mock_paths):
|
||||
"""Workspace exists — should not recreate, but still add missing templates."""
|
||||
"""Workspace exists (from fixture) — should add missing templates."""
|
||||
config_file, workspace_dir = mock_paths
|
||||
workspace_dir.mkdir(parents=True)
|
||||
config_file.write_text("{}")
|
||||
# workspace_dir already exists from fixture
|
||||
# No existing config, so onboard should proceed
|
||||
|
||||
result = runner.invoke(app, ["onboard"], input="n\n")
|
||||
result = runner.invoke(app, ["onboard"])
|
||||
|
||||
assert result.exit_code == 0
|
||||
assert "Created workspace" not in result.stdout
|
||||
assert "Created workspace" in result.stdout
|
||||
assert "Created AGENTS.md" in result.stdout
|
||||
assert (workspace_dir / "AGENTS.md").exists()
|
||||
|
||||
|
||||
+21
-28
@@ -12,15 +12,17 @@ async def test_computer_tool_screenshot():
|
||||
tool = ComputerTool20251124(vnc_host="localhost", vnc_port=5900)
|
||||
|
||||
# Mock VNC client
|
||||
with patch('nanobot.agent.tools.anthropic.computer.VNCDoToolClient') as mock_vnc:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.captureScreen = AsyncMock(return_value=b"fake_png_data")
|
||||
|
||||
# Set up async context manager
|
||||
mock_context = MagicMock()
|
||||
mock_context.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_context.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_vnc.create = MagicMock(return_value=mock_context)
|
||||
with patch('nanobot.agent.tools.anthropic.computer.vnc_api.connect') as mock_connect:
|
||||
mock_client = MagicMock()
|
||||
# Mock captureScreen to write fake PNG data to file path
|
||||
def fake_capture(path):
|
||||
from pathlib import Path
|
||||
Path(path).write_bytes(b"fake_png_data")
|
||||
mock_client.captureScreen = MagicMock(side_effect=fake_capture)
|
||||
mock_client.mouseMove = MagicMock()
|
||||
mock_client.keyPress = MagicMock()
|
||||
mock_client.refreshScreen = MagicMock()
|
||||
mock_connect.return_value = mock_client
|
||||
|
||||
result = await tool(action="screenshot")
|
||||
|
||||
@@ -34,15 +36,10 @@ async def test_computer_tool_mouse_move():
|
||||
"""Test computer tool can move mouse."""
|
||||
tool = ComputerTool20251124(vnc_host="localhost", vnc_port=5900)
|
||||
|
||||
with patch('nanobot.agent.tools.anthropic.computer.VNCDoToolClient') as mock_vnc:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.mouseMove = AsyncMock()
|
||||
|
||||
# Set up async context manager
|
||||
mock_context = MagicMock()
|
||||
mock_context.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_context.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_vnc.create = MagicMock(return_value=mock_context)
|
||||
with patch('nanobot.agent.tools.anthropic.computer.vnc_api.connect') as mock_connect:
|
||||
mock_client = MagicMock()
|
||||
mock_client.mouseMove = MagicMock()
|
||||
mock_connect.return_value = mock_client
|
||||
|
||||
result = await tool(action="mouse_move", coordinate=[100, 200])
|
||||
|
||||
@@ -56,21 +53,17 @@ async def test_computer_tool_key():
|
||||
"""Test computer tool can press keys."""
|
||||
tool = ComputerTool20251124(vnc_host="localhost", vnc_port=5900)
|
||||
|
||||
with patch('nanobot.agent.tools.anthropic.computer.VNCDoToolClient') as mock_vnc:
|
||||
mock_client = AsyncMock()
|
||||
mock_client.keyPress = AsyncMock()
|
||||
|
||||
# Set up async context manager
|
||||
mock_context = MagicMock()
|
||||
mock_context.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_context.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_vnc.create = MagicMock(return_value=mock_context)
|
||||
with patch('nanobot.agent.tools.anthropic.computer.vnc_api.connect') as mock_connect:
|
||||
mock_client = MagicMock()
|
||||
mock_client.keyPress = MagicMock()
|
||||
mock_connect.return_value = mock_client
|
||||
|
||||
result = await tool(action="key", text="Return")
|
||||
|
||||
assert isinstance(result, ToolResult)
|
||||
assert result.error is None
|
||||
mock_client.keyPress.assert_called_once_with("Return")
|
||||
# Implementation converts keys to lowercase
|
||||
mock_client.keyPress.assert_called_once_with("return")
|
||||
|
||||
|
||||
def test_computer_tool_to_params():
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Tests for config loader (get_config_path and _migrate_config)"""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from nanobot.config.loader import get_config_path, _migrate_config
|
||||
|
||||
|
||||
def test_get_config_path_default():
|
||||
"""get_config_path returns ~/.nanobot/config.json by default"""
|
||||
# Ensure NANOBOT_CONFIG is not set
|
||||
env_backup = os.environ.pop("NANOBOT_CONFIG", None)
|
||||
try:
|
||||
path = get_config_path()
|
||||
assert path == Path.home() / ".nanobot" / "config.json"
|
||||
finally:
|
||||
if env_backup:
|
||||
os.environ["NANOBOT_CONFIG"] = env_backup
|
||||
|
||||
|
||||
def test_get_config_path_with_env_var():
|
||||
"""get_config_path uses NANOBOT_CONFIG env var when set"""
|
||||
custom_path = "/tmp/test-nanobot-config.json"
|
||||
env_backup = os.environ.get("NANOBOT_CONFIG")
|
||||
try:
|
||||
os.environ["NANOBOT_CONFIG"] = custom_path
|
||||
path = get_config_path()
|
||||
assert path == Path(custom_path)
|
||||
finally:
|
||||
if env_backup:
|
||||
os.environ["NANOBOT_CONFIG"] = env_backup
|
||||
else:
|
||||
os.environ.pop("NANOBOT_CONFIG", None)
|
||||
|
||||
|
||||
def test_migrate_config_with_oauth_credentials():
|
||||
"""_migrate_config extracts api_key from oauthCredentials"""
|
||||
data = {
|
||||
"providers": {
|
||||
"anthropic": {
|
||||
"oauthCredentials": {
|
||||
"access_token": "sk-ant-test-token",
|
||||
"refresh_token": "",
|
||||
"expires_at": 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result = _migrate_config(data)
|
||||
|
||||
# api_key should be extracted
|
||||
assert result["providers"]["anthropic"]["api_key"] == "sk-ant-test-token"
|
||||
# oauthCredentials should be removed after migration
|
||||
assert "oauthCredentials" not in result["providers"]["anthropic"]
|
||||
|
||||
|
||||
def test_migrate_config_without_oauth_credentials():
|
||||
"""_migrate_config leaves config unchanged when no oauthCredentials"""
|
||||
data = {
|
||||
"providers": {
|
||||
"anthropic": {
|
||||
"api_key": "sk-ant-existing-key"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result = _migrate_config(data)
|
||||
|
||||
# Should remain unchanged
|
||||
assert result["providers"]["anthropic"]["api_key"] == "sk-ant-existing-key"
|
||||
assert "oauthCredentials" not in result["providers"]["anthropic"]
|
||||
|
||||
|
||||
def test_migrate_config_already_migrated():
|
||||
"""_migrate_config doesn't overwrite existing api_key"""
|
||||
data = {
|
||||
"providers": {
|
||||
"anthropic": {
|
||||
"api_key": "sk-ant-existing-key",
|
||||
"oauthCredentials": {
|
||||
"access_token": "sk-ant-oauth-token",
|
||||
"refresh_token": "",
|
||||
"expires_at": 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result = _migrate_config(data)
|
||||
|
||||
# Existing api_key should be preserved
|
||||
assert result["providers"]["anthropic"]["api_key"] == "sk-ant-existing-key"
|
||||
# oauthCredentials should NOT be removed (api_key already existed)
|
||||
assert "oauthCredentials" in result["providers"]["anthropic"]
|
||||
|
||||
|
||||
def test_migrate_config_empty_access_token():
|
||||
"""_migrate_config skips empty access_token"""
|
||||
data = {
|
||||
"providers": {
|
||||
"anthropic": {
|
||||
"oauthCredentials": {
|
||||
"access_token": "",
|
||||
"refresh_token": "",
|
||||
"expires_at": 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result = _migrate_config(data)
|
||||
|
||||
# api_key should not be set
|
||||
assert "api_key" not in result["providers"]["anthropic"]
|
||||
# oauthCredentials should remain (no migration happened)
|
||||
assert "oauthCredentials" in result["providers"]["anthropic"]
|
||||
|
||||
|
||||
def test_migrate_config_preserves_other_fields():
|
||||
"""_migrate_config preserves other provider config fields"""
|
||||
data = {
|
||||
"providers": {
|
||||
"anthropic": {
|
||||
"oauthCredentials": {
|
||||
"access_token": "sk-ant-test-token",
|
||||
"refresh_token": "refresh-token",
|
||||
},
|
||||
"customField": "customValue",
|
||||
"anotherField": 123,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result = _migrate_config(data)
|
||||
|
||||
# api_key added, oauthCredentials removed
|
||||
assert result["providers"]["anthropic"]["api_key"] == "sk-ant-test-token"
|
||||
assert "oauthCredentials" not in result["providers"]["anthropic"]
|
||||
# Other fields preserved
|
||||
assert result["providers"]["anthropic"]["customField"] == "customValue"
|
||||
assert result["providers"]["anthropic"]["anotherField"] == 123
|
||||
@@ -13,7 +13,7 @@ def test_oauth_token_injected_into_config(tmp_path, monkeypatch):
|
||||
# Create a minimal config file (no api key set)
|
||||
config_path = tmp_path / "config.json"
|
||||
config_path.write_text(json.dumps({
|
||||
"agents": {"defaults": {"model": "anthropic/claude-opus-4-5"}},
|
||||
"agents": {"defaults": {"model": "anthropic/claude-opus-4-7"}},
|
||||
"providers": {"anthropic": {"apiKey": ""}}
|
||||
}))
|
||||
|
||||
|
||||
@@ -1,828 +0,0 @@
|
||||
"""Test session management with cache-friendly message handling."""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from pathlib import Path
|
||||
from nanobot.session.manager import Session, SessionManager
|
||||
|
||||
# Test constants
|
||||
MEMORY_WINDOW = 50
|
||||
KEEP_COUNT = MEMORY_WINDOW // 2 # 25
|
||||
|
||||
|
||||
def create_session_with_messages(key: str, count: int, role: str = "user") -> Session:
|
||||
"""Create a session and add the specified number of messages.
|
||||
|
||||
Args:
|
||||
key: Session identifier
|
||||
count: Number of messages to add
|
||||
role: Message role (default: "user")
|
||||
|
||||
Returns:
|
||||
Session with the specified messages
|
||||
"""
|
||||
session = Session(key=key)
|
||||
for i in range(count):
|
||||
session.add_message(role, f"msg{i}")
|
||||
return session
|
||||
|
||||
|
||||
def assert_messages_content(messages: list, start_index: int, end_index: int) -> None:
|
||||
"""Assert that messages contain expected content from start to end index.
|
||||
|
||||
Args:
|
||||
messages: List of message dictionaries
|
||||
start_index: Expected first message index
|
||||
end_index: Expected last message index
|
||||
"""
|
||||
assert len(messages) > 0
|
||||
assert messages[0]["content"] == f"msg{start_index}"
|
||||
assert messages[-1]["content"] == f"msg{end_index}"
|
||||
|
||||
|
||||
def get_old_messages(session: Session, last_consolidated: int, keep_count: int) -> list:
|
||||
"""Extract messages that would be consolidated using the standard slice logic.
|
||||
|
||||
Args:
|
||||
session: The session containing messages
|
||||
last_consolidated: Index of last consolidated message
|
||||
keep_count: Number of recent messages to keep
|
||||
|
||||
Returns:
|
||||
List of messages that would be consolidated
|
||||
"""
|
||||
return session.messages[last_consolidated:-keep_count]
|
||||
|
||||
|
||||
class TestSessionLastConsolidated:
|
||||
"""Test last_consolidated tracking to avoid duplicate processing."""
|
||||
|
||||
def test_initial_last_consolidated_zero(self) -> None:
|
||||
"""Test that new session starts with last_consolidated=0."""
|
||||
session = Session(key="test:initial")
|
||||
assert session.last_consolidated == 0
|
||||
|
||||
def test_last_consolidated_persistence(self, tmp_path) -> None:
|
||||
"""Test that last_consolidated persists across save/load."""
|
||||
manager = SessionManager(Path(tmp_path))
|
||||
session1 = create_session_with_messages("test:persist", 20)
|
||||
session1.last_consolidated = 15
|
||||
manager.save(session1)
|
||||
|
||||
session2 = manager.get_or_create("test:persist")
|
||||
assert session2.last_consolidated == 15
|
||||
assert len(session2.messages) == 20
|
||||
|
||||
def test_clear_resets_last_consolidated(self) -> None:
|
||||
"""Test that clear() resets last_consolidated to 0."""
|
||||
session = create_session_with_messages("test:clear", 10)
|
||||
session.last_consolidated = 5
|
||||
|
||||
session.clear()
|
||||
assert len(session.messages) == 0
|
||||
assert session.last_consolidated == 0
|
||||
|
||||
|
||||
class TestSessionImmutableHistory:
|
||||
"""Test Session message immutability for cache efficiency."""
|
||||
|
||||
def test_initial_state(self) -> None:
|
||||
"""Test that new session has empty messages list."""
|
||||
session = Session(key="test:initial")
|
||||
assert len(session.messages) == 0
|
||||
|
||||
def test_add_messages_appends_only(self) -> None:
|
||||
"""Test that adding messages only appends, never modifies."""
|
||||
session = Session(key="test:preserve")
|
||||
session.add_message("user", "msg1")
|
||||
session.add_message("assistant", "resp1")
|
||||
session.add_message("user", "msg2")
|
||||
assert len(session.messages) == 3
|
||||
assert session.messages[0]["content"] == "msg1"
|
||||
|
||||
def test_get_history_returns_most_recent(self) -> None:
|
||||
"""Test get_history returns the most recent messages."""
|
||||
session = Session(key="test:history")
|
||||
for i in range(10):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
|
||||
history = session.get_history(max_messages=6)
|
||||
assert len(history) == 6
|
||||
assert history[0]["content"] == "msg7"
|
||||
assert history[-1]["content"] == "resp9"
|
||||
|
||||
def test_get_history_with_all_messages(self) -> None:
|
||||
"""Test get_history with max_messages larger than actual."""
|
||||
session = create_session_with_messages("test:all", 5)
|
||||
history = session.get_history(max_messages=100)
|
||||
assert len(history) == 5
|
||||
assert history[0]["content"] == "msg0"
|
||||
|
||||
def test_get_history_stable_for_same_session(self) -> None:
|
||||
"""Test that get_history returns same content for same max_messages."""
|
||||
session = create_session_with_messages("test:stable", 20)
|
||||
history1 = session.get_history(max_messages=10)
|
||||
history2 = session.get_history(max_messages=10)
|
||||
assert history1 == history2
|
||||
|
||||
def test_messages_list_never_modified(self) -> None:
|
||||
"""Test that messages list is never modified after creation."""
|
||||
session = create_session_with_messages("test:immutable", 5)
|
||||
original_len = len(session.messages)
|
||||
|
||||
session.get_history(max_messages=2)
|
||||
assert len(session.messages) == original_len
|
||||
|
||||
for _ in range(10):
|
||||
session.get_history(max_messages=3)
|
||||
assert len(session.messages) == original_len
|
||||
|
||||
|
||||
class TestSessionPersistence:
|
||||
"""Test Session persistence and reload."""
|
||||
|
||||
@pytest.fixture
|
||||
def temp_manager(self, tmp_path):
|
||||
return SessionManager(Path(tmp_path))
|
||||
|
||||
def test_persistence_roundtrip(self, temp_manager):
|
||||
"""Test that messages persist across save/load."""
|
||||
session1 = create_session_with_messages("test:persistence", 20)
|
||||
temp_manager.save(session1)
|
||||
|
||||
session2 = temp_manager.get_or_create("test:persistence")
|
||||
assert len(session2.messages) == 20
|
||||
assert session2.messages[0]["content"] == "msg0"
|
||||
assert session2.messages[-1]["content"] == "msg19"
|
||||
|
||||
def test_get_history_after_reload(self, temp_manager):
|
||||
"""Test that get_history works correctly after reload."""
|
||||
session1 = create_session_with_messages("test:reload", 30)
|
||||
temp_manager.save(session1)
|
||||
|
||||
session2 = temp_manager.get_or_create("test:reload")
|
||||
history = session2.get_history(max_messages=10)
|
||||
assert len(history) == 10
|
||||
assert history[0]["content"] == "msg20"
|
||||
assert history[-1]["content"] == "msg29"
|
||||
|
||||
def test_clear_resets_session(self, temp_manager):
|
||||
"""Test that clear() properly resets session."""
|
||||
session = create_session_with_messages("test:clear", 10)
|
||||
assert len(session.messages) == 10
|
||||
|
||||
session.clear()
|
||||
assert len(session.messages) == 0
|
||||
|
||||
|
||||
class TestConsolidationTriggerConditions:
|
||||
"""Test consolidation trigger conditions and logic."""
|
||||
|
||||
def test_consolidation_needed_when_messages_exceed_window(self):
|
||||
"""Test consolidation logic: should trigger when messages > memory_window."""
|
||||
session = create_session_with_messages("test:trigger", 60)
|
||||
|
||||
total_messages = len(session.messages)
|
||||
messages_to_process = total_messages - session.last_consolidated
|
||||
|
||||
assert total_messages > MEMORY_WINDOW
|
||||
assert messages_to_process > 0
|
||||
|
||||
expected_consolidate_count = total_messages - KEEP_COUNT
|
||||
assert expected_consolidate_count == 35
|
||||
|
||||
def test_consolidation_skipped_when_within_keep_count(self):
|
||||
"""Test consolidation skipped when total messages <= keep_count."""
|
||||
session = create_session_with_messages("test:skip", 20)
|
||||
|
||||
total_messages = len(session.messages)
|
||||
assert total_messages <= KEEP_COUNT
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
assert len(old_messages) == 0
|
||||
|
||||
def test_consolidation_skipped_when_no_new_messages(self):
|
||||
"""Test consolidation skipped when messages_to_process <= 0."""
|
||||
session = create_session_with_messages("test:already_consolidated", 40)
|
||||
session.last_consolidated = len(session.messages) - KEEP_COUNT # 15
|
||||
|
||||
# Add a few more messages
|
||||
for i in range(40, 42):
|
||||
session.add_message("user", f"msg{i}")
|
||||
|
||||
total_messages = len(session.messages)
|
||||
messages_to_process = total_messages - session.last_consolidated
|
||||
assert messages_to_process > 0
|
||||
|
||||
# Simulate last_consolidated catching up
|
||||
session.last_consolidated = total_messages - KEEP_COUNT
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
assert len(old_messages) == 0
|
||||
|
||||
|
||||
class TestLastConsolidatedEdgeCases:
|
||||
"""Test last_consolidated edge cases and data corruption scenarios."""
|
||||
|
||||
def test_last_consolidated_exceeds_message_count(self):
|
||||
"""Test behavior when last_consolidated > len(messages) (data corruption)."""
|
||||
session = create_session_with_messages("test:corruption", 10)
|
||||
session.last_consolidated = 20
|
||||
|
||||
total_messages = len(session.messages)
|
||||
messages_to_process = total_messages - session.last_consolidated
|
||||
assert messages_to_process <= 0
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, 5)
|
||||
assert len(old_messages) == 0
|
||||
|
||||
def test_last_consolidated_negative_value(self):
|
||||
"""Test behavior with negative last_consolidated (invalid state)."""
|
||||
session = create_session_with_messages("test:negative", 10)
|
||||
session.last_consolidated = -5
|
||||
|
||||
keep_count = 3
|
||||
old_messages = get_old_messages(session, session.last_consolidated, keep_count)
|
||||
|
||||
# messages[-5:-3] with 10 messages gives indices 5,6
|
||||
assert len(old_messages) == 2
|
||||
assert old_messages[0]["content"] == "msg5"
|
||||
assert old_messages[-1]["content"] == "msg6"
|
||||
|
||||
def test_messages_added_after_consolidation(self):
|
||||
"""Test correct behavior when new messages arrive after consolidation."""
|
||||
session = create_session_with_messages("test:new_messages", 40)
|
||||
session.last_consolidated = len(session.messages) - KEEP_COUNT # 15
|
||||
|
||||
# Add new messages after consolidation
|
||||
for i in range(40, 50):
|
||||
session.add_message("user", f"msg{i}")
|
||||
|
||||
total_messages = len(session.messages)
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
expected_consolidate_count = total_messages - KEEP_COUNT - session.last_consolidated
|
||||
|
||||
assert len(old_messages) == expected_consolidate_count
|
||||
assert_messages_content(old_messages, 15, 24)
|
||||
|
||||
def test_slice_behavior_when_indices_overlap(self):
|
||||
"""Test slice behavior when last_consolidated >= total - keep_count."""
|
||||
session = create_session_with_messages("test:overlap", 30)
|
||||
session.last_consolidated = 12
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, 20)
|
||||
assert len(old_messages) == 0
|
||||
|
||||
|
||||
class TestArchiveAllMode:
|
||||
"""Test archive_all mode (used by /new command)."""
|
||||
|
||||
def test_archive_all_consolidates_everything(self):
|
||||
"""Test archive_all=True consolidates all messages."""
|
||||
session = create_session_with_messages("test:archive_all", 50)
|
||||
|
||||
archive_all = True
|
||||
if archive_all:
|
||||
old_messages = session.messages
|
||||
assert len(old_messages) == 50
|
||||
|
||||
assert session.last_consolidated == 0
|
||||
|
||||
def test_archive_all_resets_last_consolidated(self):
|
||||
"""Test that archive_all mode resets last_consolidated to 0."""
|
||||
session = create_session_with_messages("test:reset", 40)
|
||||
session.last_consolidated = 15
|
||||
|
||||
archive_all = True
|
||||
if archive_all:
|
||||
session.last_consolidated = 0
|
||||
|
||||
assert session.last_consolidated == 0
|
||||
assert len(session.messages) == 40
|
||||
|
||||
def test_archive_all_vs_normal_consolidation(self):
|
||||
"""Test difference between archive_all and normal consolidation."""
|
||||
# Normal consolidation
|
||||
session1 = create_session_with_messages("test:normal", 60)
|
||||
session1.last_consolidated = len(session1.messages) - KEEP_COUNT
|
||||
|
||||
# archive_all mode
|
||||
session2 = create_session_with_messages("test:all", 60)
|
||||
session2.last_consolidated = 0
|
||||
|
||||
assert session1.last_consolidated == 35
|
||||
assert len(session1.messages) == 60
|
||||
assert session2.last_consolidated == 0
|
||||
assert len(session2.messages) == 60
|
||||
|
||||
|
||||
class TestCacheImmutability:
|
||||
"""Test that consolidation doesn't modify session.messages (cache safety)."""
|
||||
|
||||
def test_consolidation_does_not_modify_messages_list(self):
|
||||
"""Test that consolidation leaves messages list unchanged."""
|
||||
session = create_session_with_messages("test:immutable", 50)
|
||||
|
||||
original_messages = session.messages.copy()
|
||||
original_len = len(session.messages)
|
||||
session.last_consolidated = original_len - KEEP_COUNT
|
||||
|
||||
assert len(session.messages) == original_len
|
||||
assert session.messages == original_messages
|
||||
|
||||
def test_get_history_does_not_modify_messages(self):
|
||||
"""Test that get_history doesn't modify messages list."""
|
||||
session = create_session_with_messages("test:history_immutable", 40)
|
||||
original_messages = [m.copy() for m in session.messages]
|
||||
|
||||
for _ in range(5):
|
||||
history = session.get_history(max_messages=10)
|
||||
assert len(history) == 10
|
||||
|
||||
assert len(session.messages) == 40
|
||||
for i, msg in enumerate(session.messages):
|
||||
assert msg["content"] == original_messages[i]["content"]
|
||||
|
||||
def test_consolidation_only_updates_last_consolidated(self):
|
||||
"""Test that consolidation only updates last_consolidated field."""
|
||||
session = create_session_with_messages("test:field_only", 60)
|
||||
|
||||
original_messages = session.messages.copy()
|
||||
original_key = session.key
|
||||
original_metadata = session.metadata.copy()
|
||||
|
||||
session.last_consolidated = len(session.messages) - KEEP_COUNT
|
||||
|
||||
assert session.messages == original_messages
|
||||
assert session.key == original_key
|
||||
assert session.metadata == original_metadata
|
||||
assert session.last_consolidated == 35
|
||||
|
||||
|
||||
class TestSliceLogic:
|
||||
"""Test the slice logic: messages[last_consolidated:-keep_count]."""
|
||||
|
||||
def test_slice_extracts_correct_range(self):
|
||||
"""Test that slice extracts the correct message range."""
|
||||
session = create_session_with_messages("test:slice", 60)
|
||||
|
||||
old_messages = get_old_messages(session, 0, KEEP_COUNT)
|
||||
|
||||
assert len(old_messages) == 35
|
||||
assert_messages_content(old_messages, 0, 34)
|
||||
|
||||
remaining = session.messages[-KEEP_COUNT:]
|
||||
assert len(remaining) == 25
|
||||
assert_messages_content(remaining, 35, 59)
|
||||
|
||||
def test_slice_with_partial_consolidation(self):
|
||||
"""Test slice when some messages already consolidated."""
|
||||
session = create_session_with_messages("test:partial", 70)
|
||||
|
||||
last_consolidated = 30
|
||||
old_messages = get_old_messages(session, last_consolidated, KEEP_COUNT)
|
||||
|
||||
assert len(old_messages) == 15
|
||||
assert_messages_content(old_messages, 30, 44)
|
||||
|
||||
def test_slice_with_various_keep_counts(self):
|
||||
"""Test slice behavior with different keep_count values."""
|
||||
session = create_session_with_messages("test:keep_counts", 50)
|
||||
|
||||
test_cases = [(10, 40), (20, 30), (30, 20), (40, 10)]
|
||||
|
||||
for keep_count, expected_count in test_cases:
|
||||
old_messages = session.messages[0:-keep_count]
|
||||
assert len(old_messages) == expected_count
|
||||
|
||||
def test_slice_when_keep_count_exceeds_messages(self):
|
||||
"""Test slice when keep_count > len(messages)."""
|
||||
session = create_session_with_messages("test:exceed", 10)
|
||||
|
||||
old_messages = session.messages[0:-20]
|
||||
assert len(old_messages) == 0
|
||||
|
||||
|
||||
class TestEmptyAndBoundarySessions:
|
||||
"""Test empty sessions and boundary conditions."""
|
||||
|
||||
def test_empty_session_consolidation(self):
|
||||
"""Test consolidation behavior with empty session."""
|
||||
session = Session(key="test:empty")
|
||||
|
||||
assert len(session.messages) == 0
|
||||
assert session.last_consolidated == 0
|
||||
|
||||
messages_to_process = len(session.messages) - session.last_consolidated
|
||||
assert messages_to_process == 0
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
assert len(old_messages) == 0
|
||||
|
||||
def test_single_message_session(self):
|
||||
"""Test consolidation with single message."""
|
||||
session = Session(key="test:single")
|
||||
session.add_message("user", "only message")
|
||||
|
||||
assert len(session.messages) == 1
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
assert len(old_messages) == 0
|
||||
|
||||
def test_exactly_keep_count_messages(self):
|
||||
"""Test session with exactly keep_count messages."""
|
||||
session = create_session_with_messages("test:exact", KEEP_COUNT)
|
||||
|
||||
assert len(session.messages) == KEEP_COUNT
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
assert len(old_messages) == 0
|
||||
|
||||
def test_just_over_keep_count(self):
|
||||
"""Test session with one message over keep_count."""
|
||||
session = create_session_with_messages("test:over", KEEP_COUNT + 1)
|
||||
|
||||
assert len(session.messages) == 26
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
assert len(old_messages) == 1
|
||||
assert old_messages[0]["content"] == "msg0"
|
||||
|
||||
def test_very_large_session(self):
|
||||
"""Test consolidation with very large message count."""
|
||||
session = create_session_with_messages("test:large", 1000)
|
||||
|
||||
assert len(session.messages) == 1000
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
assert len(old_messages) == 975
|
||||
assert_messages_content(old_messages, 0, 974)
|
||||
|
||||
remaining = session.messages[-KEEP_COUNT:]
|
||||
assert len(remaining) == 25
|
||||
assert_messages_content(remaining, 975, 999)
|
||||
|
||||
def test_session_with_gaps_in_consolidation(self):
|
||||
"""Test session with potential gaps in consolidation history."""
|
||||
session = create_session_with_messages("test:gaps", 50)
|
||||
session.last_consolidated = 10
|
||||
|
||||
# Add more messages
|
||||
for i in range(50, 60):
|
||||
session.add_message("user", f"msg{i}")
|
||||
|
||||
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||
|
||||
expected_count = 60 - KEEP_COUNT - 10
|
||||
assert len(old_messages) == expected_count
|
||||
assert_messages_content(old_messages, 10, 34)
|
||||
|
||||
|
||||
class TestConsolidationDeduplicationGuard:
|
||||
"""Test that consolidation tasks are deduplicated and serialized."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consolidation_guard_prevents_duplicate_tasks(self, tmp_path: Path) -> None:
|
||||
"""Concurrent messages above memory_window spawn only one consolidation task."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
loop = AgentLoop(
|
||||
bus=bus, provider=provider, workspace=tmp_path, model="test-model", memory_window=10
|
||||
)
|
||||
|
||||
loop.provider.chat = AsyncMock(return_value=LLMResponse(content="ok", tool_calls=[]))
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
for i in range(15):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
loop.sessions.save(session)
|
||||
|
||||
consolidation_calls = 0
|
||||
|
||||
async def _fake_consolidate(_session, archive_all: bool = False) -> None:
|
||||
nonlocal consolidation_calls
|
||||
consolidation_calls += 1
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
loop._consolidate_memory = _fake_consolidate # type: ignore[method-assign]
|
||||
|
||||
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello")
|
||||
await loop._process_message(msg)
|
||||
await loop._process_message(msg)
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
assert consolidation_calls == 1, (
|
||||
f"Expected exactly 1 consolidation, got {consolidation_calls}"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_command_guard_prevents_concurrent_consolidation(
|
||||
self, tmp_path: Path
|
||||
) -> None:
|
||||
"""/new command does not run consolidation concurrently with in-flight consolidation."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
loop = AgentLoop(
|
||||
bus=bus, provider=provider, workspace=tmp_path, model="test-model", memory_window=10
|
||||
)
|
||||
|
||||
loop.provider.chat = AsyncMock(return_value=LLMResponse(content="ok", tool_calls=[]))
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
for i in range(15):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
loop.sessions.save(session)
|
||||
|
||||
consolidation_calls = 0
|
||||
active = 0
|
||||
max_active = 0
|
||||
|
||||
async def _fake_consolidate(_session, archive_all: bool = False) -> None:
|
||||
nonlocal consolidation_calls, active, max_active
|
||||
consolidation_calls += 1
|
||||
active += 1
|
||||
max_active = max(max_active, active)
|
||||
await asyncio.sleep(0.05)
|
||||
active -= 1
|
||||
|
||||
loop._consolidate_memory = _fake_consolidate # type: ignore[method-assign]
|
||||
|
||||
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello")
|
||||
await loop._process_message(msg)
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
await loop._process_message(new_msg)
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
assert consolidation_calls == 2, (
|
||||
f"Expected normal + /new consolidations, got {consolidation_calls}"
|
||||
)
|
||||
assert max_active == 1, (
|
||||
f"Expected serialized consolidation, observed concurrency={max_active}"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consolidation_tasks_are_referenced(self, tmp_path: Path) -> None:
|
||||
"""create_task results are tracked in _consolidation_tasks while in flight."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
loop = AgentLoop(
|
||||
bus=bus, provider=provider, workspace=tmp_path, model="test-model", memory_window=10
|
||||
)
|
||||
|
||||
loop.provider.chat = AsyncMock(return_value=LLMResponse(content="ok", tool_calls=[]))
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
for i in range(15):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
loop.sessions.save(session)
|
||||
|
||||
started = asyncio.Event()
|
||||
|
||||
async def _slow_consolidate(_session, archive_all: bool = False) -> None:
|
||||
started.set()
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
loop._consolidate_memory = _slow_consolidate # type: ignore[method-assign]
|
||||
|
||||
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello")
|
||||
await loop._process_message(msg)
|
||||
|
||||
await started.wait()
|
||||
assert len(loop._consolidation_tasks) == 1, "Task must be referenced while in-flight"
|
||||
|
||||
await asyncio.sleep(0.15)
|
||||
assert len(loop._consolidation_tasks) == 0, (
|
||||
"Task reference must be removed after completion"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_waits_for_inflight_consolidation_and_preserves_messages(
|
||||
self, tmp_path: Path
|
||||
) -> None:
|
||||
"""/new waits for in-flight consolidation and archives before clear."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
loop = AgentLoop(
|
||||
bus=bus, provider=provider, workspace=tmp_path, model="test-model", memory_window=10
|
||||
)
|
||||
|
||||
loop.provider.chat = AsyncMock(return_value=LLMResponse(content="ok", tool_calls=[]))
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
for i in range(15):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
loop.sessions.save(session)
|
||||
|
||||
started = asyncio.Event()
|
||||
release = asyncio.Event()
|
||||
archived_count = 0
|
||||
|
||||
async def _fake_consolidate(sess, archive_all: bool = False) -> bool:
|
||||
nonlocal archived_count
|
||||
if archive_all:
|
||||
archived_count = len(sess.messages)
|
||||
return True
|
||||
started.set()
|
||||
await release.wait()
|
||||
return True
|
||||
|
||||
loop._consolidate_memory = _fake_consolidate # type: ignore[method-assign]
|
||||
|
||||
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello")
|
||||
await loop._process_message(msg)
|
||||
await started.wait()
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
pending_new = asyncio.create_task(loop._process_message(new_msg))
|
||||
|
||||
await asyncio.sleep(0.02)
|
||||
assert not pending_new.done(), "/new should wait while consolidation is in-flight"
|
||||
|
||||
release.set()
|
||||
response = await pending_new
|
||||
assert response is not None
|
||||
assert "new session started" in response.content.lower()
|
||||
assert archived_count > 0, "Expected /new archival to process a non-empty snapshot"
|
||||
|
||||
session_after = loop.sessions.get_or_create("cli:test")
|
||||
assert session_after.messages == [], "Session should be cleared after successful archival"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_does_not_clear_session_when_archive_fails(self, tmp_path: Path) -> None:
|
||||
"""/new must keep session data if archive step reports failure."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
loop = AgentLoop(
|
||||
bus=bus, provider=provider, workspace=tmp_path, model="test-model", memory_window=10
|
||||
)
|
||||
|
||||
loop.provider.chat = AsyncMock(return_value=LLMResponse(content="ok", tool_calls=[]))
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
for i in range(5):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
loop.sessions.save(session)
|
||||
before_count = len(session.messages)
|
||||
|
||||
async def _failing_consolidate(sess, archive_all: bool = False) -> bool:
|
||||
if archive_all:
|
||||
return False
|
||||
return True
|
||||
|
||||
loop._consolidate_memory = _failing_consolidate # type: ignore[method-assign]
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
response = await loop._process_message(new_msg)
|
||||
|
||||
assert response is not None
|
||||
assert "failed" in response.content.lower()
|
||||
session_after = loop.sessions.get_or_create("cli:test")
|
||||
assert len(session_after.messages) == before_count, (
|
||||
"Session must remain intact when /new archival fails"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_archives_only_unconsolidated_messages_after_inflight_task(
|
||||
self, tmp_path: Path
|
||||
) -> None:
|
||||
"""/new should archive only messages not yet consolidated by prior task."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
loop = AgentLoop(
|
||||
bus=bus, provider=provider, workspace=tmp_path, model="test-model", memory_window=10
|
||||
)
|
||||
|
||||
loop.provider.chat = AsyncMock(return_value=LLMResponse(content="ok", tool_calls=[]))
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
for i in range(15):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
loop.sessions.save(session)
|
||||
|
||||
started = asyncio.Event()
|
||||
release = asyncio.Event()
|
||||
archived_count = -1
|
||||
|
||||
async def _fake_consolidate(sess, archive_all: bool = False) -> bool:
|
||||
nonlocal archived_count
|
||||
if archive_all:
|
||||
archived_count = len(sess.messages)
|
||||
return True
|
||||
|
||||
started.set()
|
||||
await release.wait()
|
||||
sess.last_consolidated = len(sess.messages) - 3
|
||||
return True
|
||||
|
||||
loop._consolidate_memory = _fake_consolidate # type: ignore[method-assign]
|
||||
|
||||
msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello")
|
||||
await loop._process_message(msg)
|
||||
await started.wait()
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
pending_new = asyncio.create_task(loop._process_message(new_msg))
|
||||
await asyncio.sleep(0.02)
|
||||
assert not pending_new.done()
|
||||
|
||||
release.set()
|
||||
response = await pending_new
|
||||
|
||||
assert response is not None
|
||||
assert "new session started" in response.content.lower()
|
||||
assert archived_count == 3, (
|
||||
f"Expected only unconsolidated tail to archive, got {archived_count}"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_cleans_up_consolidation_lock_for_invalidated_session(
|
||||
self, tmp_path: Path
|
||||
) -> None:
|
||||
"""/new should remove lock entry for fully invalidated session key."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
loop = AgentLoop(
|
||||
bus=bus, provider=provider, workspace=tmp_path, model="test-model", memory_window=10
|
||||
)
|
||||
|
||||
loop.provider.chat = AsyncMock(return_value=LLMResponse(content="ok", tool_calls=[]))
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
for i in range(3):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
loop.sessions.save(session)
|
||||
|
||||
# Ensure lock exists before /new.
|
||||
loop._consolidation_locks.setdefault(session.key, asyncio.Lock())
|
||||
assert session.key in loop._consolidation_locks
|
||||
|
||||
async def _ok_consolidate(sess, archive_all: bool = False) -> bool:
|
||||
return True
|
||||
|
||||
loop._consolidate_memory = _ok_consolidate # type: ignore[method-assign]
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
response = await loop._process_message(new_msg)
|
||||
|
||||
assert response is not None
|
||||
assert "new session started" in response.content.lower()
|
||||
assert session.key not in loop._consolidation_locks
|
||||
@@ -40,7 +40,7 @@ def test_system_prompt_stays_stable_when_clock_changes(tmp_path, monkeypatch) ->
|
||||
|
||||
|
||||
def test_runtime_context_is_separate_untrusted_user_message(tmp_path) -> None:
|
||||
"""Runtime metadata should be a separate user message before the actual user message."""
|
||||
"""Runtime metadata should be included in the system prompt."""
|
||||
workspace = _make_workspace(tmp_path)
|
||||
builder = ContextBuilder(workspace)
|
||||
|
||||
@@ -51,16 +51,12 @@ def test_runtime_context_is_separate_untrusted_user_message(tmp_path) -> None:
|
||||
chat_id="direct",
|
||||
)
|
||||
|
||||
# Runtime context should be in the system prompt
|
||||
assert messages[0]["role"] == "system"
|
||||
assert "## Current Session" not in messages[0]["content"]
|
||||
|
||||
assert messages[-2]["role"] == "user"
|
||||
runtime_content = messages[-2]["content"]
|
||||
assert isinstance(runtime_content, str)
|
||||
assert ContextBuilder._RUNTIME_CONTEXT_TAG in runtime_content
|
||||
assert "Current Time:" in runtime_content
|
||||
assert "Channel: cli" in runtime_content
|
||||
assert "Chat ID: direct" in runtime_content
|
||||
assert "## Current Session" in messages[0]["content"]
|
||||
assert "Channel: cli" in messages[0]["content"]
|
||||
assert "Chat ID: direct" in messages[0]["content"]
|
||||
|
||||
# The actual user message should be the last message
|
||||
assert messages[-1]["role"] == "user"
|
||||
assert messages[-1]["content"] == "Return exactly: OK"
|
||||
|
||||
@@ -113,4 +113,4 @@ def test_edit_tool_to_params():
|
||||
params = tool.to_params()
|
||||
|
||||
assert params["type"] == "text_editor_20250728"
|
||||
assert params["name"] == "str_replace_editor"
|
||||
assert params["name"] == "str_replace_based_edit_tool"
|
||||
|
||||
@@ -3,27 +3,12 @@ import asyncio
|
||||
import pytest
|
||||
|
||||
from nanobot.heartbeat.service import HeartbeatService
|
||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||
|
||||
|
||||
class DummyProvider:
|
||||
def __init__(self, responses: list[LLMResponse]):
|
||||
self._responses = list(responses)
|
||||
|
||||
async def chat(self, *args, **kwargs) -> LLMResponse:
|
||||
if self._responses:
|
||||
return self._responses.pop(0)
|
||||
return LLMResponse(content="", tool_calls=[])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_is_idempotent(tmp_path) -> None:
|
||||
provider = DummyProvider([])
|
||||
|
||||
service = HeartbeatService(
|
||||
workspace=tmp_path,
|
||||
provider=provider,
|
||||
model="openai/gpt-4o-mini",
|
||||
interval_s=9999,
|
||||
enabled=True,
|
||||
)
|
||||
@@ -38,80 +23,36 @@ async def test_start_is_idempotent(tmp_path) -> None:
|
||||
await asyncio.sleep(0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decide_returns_skip_when_no_tool_call(tmp_path) -> None:
|
||||
provider = DummyProvider([LLMResponse(content="no tool call", tool_calls=[])])
|
||||
service = HeartbeatService(
|
||||
workspace=tmp_path,
|
||||
provider=provider,
|
||||
model="openai/gpt-4o-mini",
|
||||
)
|
||||
|
||||
action, tasks = await service._decide("heartbeat content")
|
||||
assert action == "skip"
|
||||
assert tasks == ""
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trigger_now_executes_when_decision_is_run(tmp_path) -> None:
|
||||
(tmp_path / "HEARTBEAT.md").write_text("- [ ] do thing", encoding="utf-8")
|
||||
|
||||
provider = DummyProvider([
|
||||
LLMResponse(
|
||||
content="",
|
||||
tool_calls=[
|
||||
ToolCallRequest(
|
||||
id="hb_1",
|
||||
name="heartbeat",
|
||||
arguments={"action": "run", "tasks": "check open tasks"},
|
||||
)
|
||||
],
|
||||
)
|
||||
])
|
||||
called_with: list[tuple[str, dict | None]] = []
|
||||
|
||||
called_with: list[str] = []
|
||||
|
||||
async def _on_execute(tasks: str) -> str:
|
||||
called_with.append(tasks)
|
||||
async def _on_heartbeat(prompt: str, metadata: dict | None = None) -> str:
|
||||
called_with.append((prompt, metadata))
|
||||
return "done"
|
||||
|
||||
service = HeartbeatService(
|
||||
workspace=tmp_path,
|
||||
provider=provider,
|
||||
model="openai/gpt-4o-mini",
|
||||
on_execute=_on_execute,
|
||||
on_heartbeat=_on_heartbeat,
|
||||
)
|
||||
|
||||
result = await service.trigger_now()
|
||||
assert result == "done"
|
||||
assert called_with == ["check open tasks"]
|
||||
assert len(called_with) == 1
|
||||
prompt, metadata = called_with[0]
|
||||
assert "HEARTBEAT.md" in prompt
|
||||
assert metadata == {"suppress_output": True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trigger_now_returns_none_when_decision_is_skip(tmp_path) -> None:
|
||||
async def test_trigger_now_returns_none_when_no_callback(tmp_path) -> None:
|
||||
(tmp_path / "HEARTBEAT.md").write_text("- [ ] do thing", encoding="utf-8")
|
||||
|
||||
provider = DummyProvider([
|
||||
LLMResponse(
|
||||
content="",
|
||||
tool_calls=[
|
||||
ToolCallRequest(
|
||||
id="hb_1",
|
||||
name="heartbeat",
|
||||
arguments={"action": "skip"},
|
||||
)
|
||||
],
|
||||
)
|
||||
])
|
||||
|
||||
async def _on_execute(tasks: str) -> str:
|
||||
return tasks
|
||||
|
||||
service = HeartbeatService(
|
||||
workspace=tmp_path,
|
||||
provider=provider,
|
||||
model="openai/gpt-4o-mini",
|
||||
on_execute=_on_execute,
|
||||
on_heartbeat=None, # No callback
|
||||
)
|
||||
|
||||
assert await service.trigger_now() is None
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
"""Test auto-consolidation on long context 429 errors."""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from nanobot.providers.base import LongContextError, LLMResponse
|
||||
|
||||
|
||||
def test_long_context_error_is_exception():
|
||||
"""LongContextError should be a distinct exception class."""
|
||||
err = LongContextError("too long")
|
||||
assert isinstance(err, Exception)
|
||||
assert str(err) == "too long"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_raises_long_context_error_on_long_context_429():
|
||||
"""Provider should raise LongContextError immediately for long-context 429s."""
|
||||
from nanobot.providers.anthropic_oauth import AnthropicOAuthProvider
|
||||
|
||||
provider = AnthropicOAuthProvider(
|
||||
oauth_token="sk-ant-oat01-test-token",
|
||||
default_model="claude-sonnet-4-6",
|
||||
)
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 429
|
||||
mock_response.text = '{"type":"error","error":{"type":"rate_limit_error","message":"Extra usage is required for long context requests."}}'
|
||||
mock_response.headers = {}
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = mock_response
|
||||
provider._client = mock_client
|
||||
|
||||
with pytest.raises(LongContextError, match="Context too long"):
|
||||
await provider._make_request(
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
)
|
||||
|
||||
# Should NOT retry — only one call
|
||||
assert mock_client.post.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_retries_normal_429():
|
||||
"""Provider should still retry normal 429s (not long-context)."""
|
||||
from nanobot.providers.anthropic_oauth import AnthropicOAuthProvider
|
||||
|
||||
provider = AnthropicOAuthProvider(
|
||||
oauth_token="sk-ant-oat01-test-token",
|
||||
default_model="claude-sonnet-4-6",
|
||||
)
|
||||
|
||||
rate_limit_response = MagicMock()
|
||||
rate_limit_response.status_code = 429
|
||||
rate_limit_response.text = '{"type":"error","error":{"type":"rate_limit_error","message":"Rate limit exceeded"}}'
|
||||
rate_limit_response.headers = {}
|
||||
|
||||
success_response = MagicMock()
|
||||
success_response.status_code = 200
|
||||
success_response.headers = {}
|
||||
success_response.json.return_value = {
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1},
|
||||
}
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.side_effect = [rate_limit_response, success_response]
|
||||
provider._client = mock_client
|
||||
|
||||
result = await provider._make_request(
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
)
|
||||
|
||||
# Should have retried and succeeded
|
||||
assert mock_client.post.call_count == 2
|
||||
assert result["stop_reason"] == "end_turn"
|
||||
@@ -676,7 +676,7 @@ async def test_on_media_message_respects_declared_size_limit(
|
||||
assert client.download_calls == []
|
||||
assert len(handled) == 1
|
||||
assert handled[0]["media"] == []
|
||||
assert handled[0]["metadata"]["attachments"] == []
|
||||
assert handled[0]["metadata"].get("attachments", []) == []
|
||||
assert "[attachment: large.bin - too large]" in handled[0]["content"]
|
||||
|
||||
|
||||
@@ -712,7 +712,7 @@ async def test_on_media_message_uses_server_limit_when_smaller_than_local_limit(
|
||||
assert client.download_calls == []
|
||||
assert len(handled) == 1
|
||||
assert handled[0]["media"] == []
|
||||
assert handled[0]["metadata"]["attachments"] == []
|
||||
assert handled[0]["metadata"].get("attachments", []) == []
|
||||
assert "[attachment: large.bin - too large]" in handled[0]["content"]
|
||||
|
||||
|
||||
@@ -746,7 +746,7 @@ async def test_on_media_message_handles_download_error(monkeypatch, tmp_path) ->
|
||||
assert len(client.download_calls) == 1
|
||||
assert len(handled) == 1
|
||||
assert handled[0]["media"] == []
|
||||
assert handled[0]["metadata"]["attachments"] == []
|
||||
assert handled[0]["metadata"].get("attachments", []) == []
|
||||
assert "[attachment: photo.png - download failed]" in handled[0]["content"]
|
||||
|
||||
|
||||
@@ -830,7 +830,7 @@ async def test_on_media_message_handles_decrypt_error(monkeypatch, tmp_path) ->
|
||||
|
||||
assert len(handled) == 1
|
||||
assert handled[0]["media"] == []
|
||||
assert handled[0]["metadata"]["attachments"] == []
|
||||
assert handled[0]["metadata"].get("attachments", []) == []
|
||||
assert "[attachment: secret.txt - download failed]" in handled[0]["content"]
|
||||
|
||||
|
||||
@@ -972,7 +972,6 @@ async def test_send_passes_thread_relates_to_to_attachment_upload(monkeypatch) -
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
async def _fake_upload_and_send_attachment(
|
||||
*,
|
||||
room_id: str,
|
||||
path: Path,
|
||||
limit_bytes: int,
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
"""Test mem0 fact extraction calls provider with thinking disabled."""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from pathlib import Path
|
||||
|
||||
from nanobot.providers.base import LLMResponse
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_provider():
|
||||
provider = AsyncMock()
|
||||
provider.chat = AsyncMock(return_value=LLMResponse(
|
||||
content='{"facts": ["user likes Python", "user works on nanobot"]}',
|
||||
finish_reason="end_turn",
|
||||
))
|
||||
return provider
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mem0_store(tmp_path):
|
||||
"""Create a Mem0MemoryStore with mocked mem0 dependency."""
|
||||
# We can't import Mem0MemoryStore at module level because it requires
|
||||
# the mem0 package. Instead, we test extract_facts as a standalone method
|
||||
# by constructing a minimal instance.
|
||||
try:
|
||||
from nanobot.agent.memory_mem0 import Mem0MemoryStore
|
||||
store = Mem0MemoryStore(workspace=tmp_path)
|
||||
return store
|
||||
except ImportError:
|
||||
pytest.skip("mem0 not installed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_facts_passes_thinking_budget_zero(mock_provider):
|
||||
"""extract_facts must pass thinking_budget=0 to provider.chat().
|
||||
|
||||
Without this, the provider inherits its instance default (e.g. 10000),
|
||||
causing the model to spend tokens on thinking instead of outputting JSON.
|
||||
"""
|
||||
try:
|
||||
from nanobot.agent.memory_mem0 import Mem0MemoryStore
|
||||
except ImportError:
|
||||
pytest.skip("mem0 not installed")
|
||||
|
||||
# Create a minimal instance without full mem0 init
|
||||
store = object.__new__(Mem0MemoryStore)
|
||||
store.custom_prompt = "Extract facts as JSON: "
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "I like Python programming"},
|
||||
{"role": "assistant", "content": "That's great! Python is versatile."},
|
||||
]
|
||||
|
||||
facts = await store.extract_facts(messages, mock_provider, "claude-sonnet-4-6")
|
||||
|
||||
# Verify provider.chat was called with thinking_budget=0
|
||||
mock_provider.chat.assert_called_once()
|
||||
call_kwargs = mock_provider.chat.call_args.kwargs
|
||||
assert call_kwargs["thinking_budget"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_facts_returns_parsed_facts(mock_provider):
|
||||
"""extract_facts should parse JSON response into a list of fact strings."""
|
||||
try:
|
||||
from nanobot.agent.memory_mem0 import Mem0MemoryStore
|
||||
except ImportError:
|
||||
pytest.skip("mem0 not installed")
|
||||
|
||||
store = object.__new__(Mem0MemoryStore)
|
||||
store.custom_prompt = "Extract facts as JSON: "
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "I like Python programming"},
|
||||
]
|
||||
|
||||
facts = await store.extract_facts(messages, mock_provider, "claude-sonnet-4-6")
|
||||
|
||||
assert facts == ["user likes Python", "user works on nanobot"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_facts_handles_empty_response():
|
||||
"""extract_facts should return empty list when provider returns no content."""
|
||||
try:
|
||||
from nanobot.agent.memory_mem0 import Mem0MemoryStore
|
||||
except ImportError:
|
||||
pytest.skip("mem0 not installed")
|
||||
|
||||
provider = AsyncMock()
|
||||
provider.chat = AsyncMock(return_value=LLMResponse(
|
||||
content="",
|
||||
finish_reason="end_turn",
|
||||
))
|
||||
|
||||
store = object.__new__(Mem0MemoryStore)
|
||||
store.custom_prompt = "Extract facts as JSON: "
|
||||
|
||||
messages = [{"role": "user", "content": "Hello there"}]
|
||||
|
||||
facts = await store.extract_facts(messages, provider, "claude-sonnet-4-6")
|
||||
|
||||
assert facts == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_facts_skips_empty_messages():
|
||||
"""extract_facts should return empty list when all messages have empty content."""
|
||||
try:
|
||||
from nanobot.agent.memory_mem0 import Mem0MemoryStore
|
||||
except ImportError:
|
||||
pytest.skip("mem0 not installed")
|
||||
|
||||
provider = AsyncMock()
|
||||
|
||||
store = object.__new__(Mem0MemoryStore)
|
||||
store.custom_prompt = "Extract facts as JSON: "
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": ""},
|
||||
{"role": "assistant", "content": ""},
|
||||
]
|
||||
|
||||
facts = await store.extract_facts(messages, provider, "claude-sonnet-4-6")
|
||||
|
||||
assert facts == []
|
||||
# Provider should not be called when there's no content
|
||||
provider.chat.assert_not_called()
|
||||
@@ -0,0 +1,153 @@
|
||||
"""Tests for message visibility signing (hidden intermediate messages)."""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from nanobot.agent.context import ContextBuilder
|
||||
from nanobot.agent.visibility import compute_signature, sign_content
|
||||
from nanobot.session.manager import Session
|
||||
|
||||
|
||||
class TestComputeSignature:
|
||||
"""Tests for compute_signature()."""
|
||||
|
||||
def test_returns_8_char_hex(self):
|
||||
sig = compute_signature("hello")
|
||||
assert len(sig) == 8
|
||||
assert all(c in "0123456789abcdef" for c in sig)
|
||||
|
||||
def test_deterministic(self):
|
||||
assert compute_signature("hello") == compute_signature("hello")
|
||||
|
||||
def test_different_content_different_sig(self):
|
||||
assert compute_signature("hello") != compute_signature("world")
|
||||
|
||||
def test_sign_content_uses_compute_signature(self):
|
||||
"""sign_content should produce [HIDDEN:{compute_signature(content)}] prefix."""
|
||||
content = "test message"
|
||||
sig = compute_signature(content)
|
||||
assert sign_content(content) == f"[HIDDEN:{sig}] {content}"
|
||||
|
||||
|
||||
class TestAddAssistantMessage:
|
||||
"""Tests for _hidden_sig in add_assistant_message()."""
|
||||
|
||||
def setup_method(self):
|
||||
self.ctx = ContextBuilder(Path("/tmp"))
|
||||
|
||||
def test_intermediate_message_gets_hidden_sig(self):
|
||||
msgs: list = []
|
||||
tool_calls = [{"id": "tc1", "type": "function", "function": {"name": "test", "arguments": "{}"}}]
|
||||
self.ctx.add_assistant_message(msgs, "thinking...", tool_calls)
|
||||
|
||||
assert msgs[0].get("_hidden_sig") is not None
|
||||
assert msgs[0]["_hidden_sig"] == compute_signature("thinking...")
|
||||
|
||||
def test_final_message_no_hidden_sig(self):
|
||||
msgs: list = []
|
||||
self.ctx.add_assistant_message(msgs, "Here is the answer", None)
|
||||
|
||||
assert "_hidden_sig" not in msgs[0]
|
||||
|
||||
def test_empty_content_signed(self):
|
||||
msgs: list = []
|
||||
tool_calls = [{"id": "tc1", "type": "function", "function": {"name": "test", "arguments": "{}"}}]
|
||||
self.ctx.add_assistant_message(msgs, None, tool_calls)
|
||||
|
||||
assert msgs[0]["_hidden_sig"] == compute_signature("")
|
||||
|
||||
|
||||
class TestAddToolResult:
|
||||
"""Tests for _hidden_sig in add_tool_result()."""
|
||||
|
||||
def setup_method(self):
|
||||
self.ctx = ContextBuilder(Path("/tmp"))
|
||||
|
||||
def test_tool_result_gets_hidden_sig(self):
|
||||
msgs: list = []
|
||||
self.ctx.add_tool_result(msgs, "tc1", "read_file", "file contents here")
|
||||
|
||||
assert msgs[0]["_hidden_sig"] == compute_signature("file contents here")
|
||||
|
||||
def test_tool_result_non_string_content(self):
|
||||
msgs: list = []
|
||||
# Multipart content (e.g. image) is a list, not a string
|
||||
self.ctx.add_tool_result(msgs, "tc1", "screenshot", [{"type": "text", "text": "ok"}])
|
||||
|
||||
assert msgs[0]["_hidden_sig"] == compute_signature("")
|
||||
|
||||
|
||||
class TestGetHistoryPrefix:
|
||||
"""Tests for get_history() applying [HIDDEN:sig] prefix."""
|
||||
|
||||
def test_hidden_sig_applied_at_read_time(self):
|
||||
session = Session(key="test")
|
||||
sig = compute_signature("thinking...")
|
||||
session.messages = [
|
||||
{"role": "assistant", "content": "thinking...", "tool_calls": [{}], "_hidden_sig": sig},
|
||||
]
|
||||
|
||||
history = session.get_history()
|
||||
assert history[0]["content"] == f"[HIDDEN:{sig}] thinking..."
|
||||
assert "_hidden_sig" not in history[0]
|
||||
|
||||
def test_no_prefix_without_hidden_sig(self):
|
||||
session = Session(key="test")
|
||||
session.messages = [
|
||||
{"role": "assistant", "content": "Here is the answer"},
|
||||
]
|
||||
|
||||
history = session.get_history()
|
||||
assert history[0]["content"] == "Here is the answer"
|
||||
|
||||
def test_tool_result_gets_prefix(self):
|
||||
session = Session(key="test")
|
||||
sig = compute_signature("file contents")
|
||||
session.messages = [
|
||||
{"role": "tool", "tool_call_id": "tc1", "name": "read", "content": "file contents", "_hidden_sig": sig},
|
||||
]
|
||||
|
||||
history = session.get_history()
|
||||
assert history[0]["content"] == f"[HIDDEN:{sig}] file contents"
|
||||
|
||||
def test_roundtrip_jsonl(self, tmp_path):
|
||||
"""Write to session JSONL, reload, verify get_history() produces correct prefix."""
|
||||
from nanobot.session.manager import SessionManager
|
||||
|
||||
workspace = tmp_path / "workspace"
|
||||
workspace.mkdir()
|
||||
mgr = SessionManager(workspace)
|
||||
|
||||
session = mgr.get_or_create("test:roundtrip")
|
||||
sig = compute_signature("intermediate")
|
||||
session.add_raw_message({
|
||||
"role": "assistant",
|
||||
"content": "intermediate",
|
||||
"tool_calls": [{"id": "tc1", "type": "function", "function": {"name": "x", "arguments": "{}"}}],
|
||||
"_hidden_sig": sig,
|
||||
})
|
||||
session.add_raw_message({
|
||||
"role": "assistant",
|
||||
"content": "final answer",
|
||||
})
|
||||
mgr.save(session)
|
||||
|
||||
# Reload from disk
|
||||
mgr.invalidate("test:roundtrip")
|
||||
reloaded = mgr.get_or_create("test:roundtrip")
|
||||
history = reloaded.get_history()
|
||||
|
||||
assert history[0]["content"] == f"[HIDDEN:{sig}] intermediate"
|
||||
assert history[1]["content"] == "final answer"
|
||||
|
||||
def test_idempotent_across_calls(self):
|
||||
"""Same prefix produced every call (cache stability)."""
|
||||
session = Session(key="test")
|
||||
sig = compute_signature("msg")
|
||||
session.messages = [
|
||||
{"role": "assistant", "content": "msg", "_hidden_sig": sig},
|
||||
]
|
||||
|
||||
h1 = session.get_history()
|
||||
h2 = session.get_history()
|
||||
assert h1[0]["content"] == h2[0]["content"]
|
||||
@@ -43,15 +43,12 @@ def test_native_tools_registered(mock_provider, mock_bus, tmp_path):
|
||||
|
||||
# Verify native tools are registered (using their internal names)
|
||||
assert "bash" in tool_names, "bash tool should be registered"
|
||||
assert "str_replace_editor" in tool_names, "str_replace_editor tool should be registered"
|
||||
assert "computer" in tool_names, "computer tool should be registered"
|
||||
assert "str_replace_based_edit_tool" in tool_names, "str_replace_based_edit_tool tool should be registered"
|
||||
# Note: computer tool is intentionally disabled by default (requires VNC setup)
|
||||
|
||||
# Verify we can get the tool instances
|
||||
bash_tool = loop.tools.get("bash")
|
||||
assert isinstance(bash_tool, BashTool20250124)
|
||||
|
||||
editor_tool = loop.tools.get("str_replace_editor")
|
||||
editor_tool = loop.tools.get("str_replace_based_edit_tool")
|
||||
assert isinstance(editor_tool, EditTool20250728)
|
||||
|
||||
computer_tool = loop.tools.get("computer")
|
||||
assert isinstance(computer_tool, ComputerTool20251124)
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
"""Test that the Anthropic OAuth identity block is always included in API requests."""
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
|
||||
import httpx
|
||||
|
||||
from nanobot.providers.anthropic_oauth import AnthropicOAuthProvider
|
||||
from nanobot.providers.oauth_utils import get_claude_code_system_prefix
|
||||
|
||||
|
||||
IDENTITY_TEXT = get_claude_code_system_prefix()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def provider():
|
||||
return AnthropicOAuthProvider(
|
||||
oauth_token="sk-ant-oat01-test-token",
|
||||
default_model="claude-opus-4-7",
|
||||
)
|
||||
|
||||
|
||||
def _mock_response(status_code=200, json_data=None):
|
||||
"""Create a mock httpx.Response."""
|
||||
resp = MagicMock(spec=httpx.Response)
|
||||
resp.status_code = status_code
|
||||
resp.headers = {}
|
||||
resp.text = ""
|
||||
resp.json.return_value = json_data or {
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1},
|
||||
}
|
||||
return resp
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_identity_block_present_with_system_prompt(provider):
|
||||
"""When a system prompt is provided, identity block is the first system block."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = _mock_response()
|
||||
provider._client = mock_client
|
||||
|
||||
await provider._make_request(
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
system="You are a helpful assistant.",
|
||||
)
|
||||
|
||||
call_kwargs = mock_client.post.call_args
|
||||
payload = call_kwargs.kwargs["json"] if "json" in call_kwargs.kwargs else call_kwargs[1]["json"]
|
||||
system_blocks = payload["system"]
|
||||
|
||||
assert len(system_blocks) == 2
|
||||
assert system_blocks[0]["type"] == "text"
|
||||
assert system_blocks[0]["text"] == IDENTITY_TEXT
|
||||
assert system_blocks[1]["text"] == "You are a helpful assistant."
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_identity_block_present_without_system_prompt(provider):
|
||||
"""When no system prompt is provided, identity block is still included.
|
||||
|
||||
This is the critical fix: extract_facts and similar calls pass system=None,
|
||||
but Anthropic requires the identity block for OAuth tokens.
|
||||
"""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = _mock_response()
|
||||
provider._client = mock_client
|
||||
|
||||
await provider._make_request(
|
||||
messages=[{"role": "user", "content": "extract facts"}],
|
||||
system=None,
|
||||
)
|
||||
|
||||
call_kwargs = mock_client.post.call_args
|
||||
payload = call_kwargs.kwargs["json"] if "json" in call_kwargs.kwargs else call_kwargs[1]["json"]
|
||||
system_blocks = payload["system"]
|
||||
|
||||
assert len(system_blocks) == 1
|
||||
assert system_blocks[0]["type"] == "text"
|
||||
assert system_blocks[0]["text"] == IDENTITY_TEXT
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_identity_block_present_with_empty_string_system(provider):
|
||||
"""Empty string system prompt should still include the identity block."""
|
||||
mock_client = AsyncMock()
|
||||
mock_client.post.return_value = _mock_response()
|
||||
provider._client = mock_client
|
||||
|
||||
await provider._make_request(
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
system="",
|
||||
)
|
||||
|
||||
call_kwargs = mock_client.post.call_args
|
||||
payload = call_kwargs.kwargs["json"] if "json" in call_kwargs.kwargs else call_kwargs[1]["json"]
|
||||
system_blocks = payload["system"]
|
||||
|
||||
# Empty string is falsy, so should go through the else branch
|
||||
assert len(system_blocks) == 1
|
||||
assert system_blocks[0]["text"] == IDENTITY_TEXT
|
||||
@@ -17,7 +17,7 @@ def test_get_auth_headers_oauth():
|
||||
assert "Authorization" in headers
|
||||
assert headers["Authorization"] == "Bearer sk-ant-oat01-xxx"
|
||||
assert "x-api-key" not in headers
|
||||
assert headers["anthropic-beta"] == "claude-code-20250219,oauth-2025-04-20"
|
||||
assert headers["anthropic-beta"] == "claude-code-20250219,oauth-2025-04-20,context-management-2025-06-27"
|
||||
|
||||
|
||||
def test_get_auth_headers_api_key():
|
||||
|
||||
@@ -9,7 +9,7 @@ def test_create_provider_oauth_token():
|
||||
"""OAuth tokens should create AnthropicOAuthProvider."""
|
||||
provider = create_provider(
|
||||
api_key="sk-ant-oat01-test-token",
|
||||
model="anthropic/claude-opus-4-5"
|
||||
model="anthropic/claude-opus-4-7"
|
||||
)
|
||||
assert isinstance(provider, AnthropicOAuthProvider)
|
||||
|
||||
@@ -18,7 +18,7 @@ def test_create_provider_regular_key():
|
||||
"""Regular API keys should create LiteLLMProvider."""
|
||||
provider = create_provider(
|
||||
api_key="sk-ant-api03-regular-key",
|
||||
model="anthropic/claude-opus-4-5"
|
||||
model="anthropic/claude-opus-4-7"
|
||||
)
|
||||
assert isinstance(provider, LiteLLMProvider)
|
||||
|
||||
@@ -27,6 +27,6 @@ def test_create_provider_openrouter():
|
||||
"""OpenRouter keys should create LiteLLMProvider."""
|
||||
provider = create_provider(
|
||||
api_key="sk-or-v1-xxx",
|
||||
model="anthropic/claude-opus-4-5"
|
||||
model="anthropic/claude-opus-4-7"
|
||||
)
|
||||
assert isinstance(provider, LiteLLMProvider)
|
||||
|
||||
@@ -31,7 +31,7 @@ async def test_registry_executes_edit_tool():
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
test_file = str(Path(tmpdir) / "test.txt")
|
||||
|
||||
result = await registry.execute("str_replace_editor", {
|
||||
result = await registry.execute("str_replace_based_edit_tool", {
|
||||
"command": "create",
|
||||
"path": test_file,
|
||||
"file_text": "Hello, world!"
|
||||
|
||||
@@ -5,14 +5,14 @@ from nanobot.providers.registry import should_use_oauth_provider
|
||||
|
||||
def test_should_use_oauth_for_oat_token():
|
||||
"""OAuth provider should be used for sk-ant-oat tokens."""
|
||||
assert should_use_oauth_provider("sk-ant-oat01-xxx", "anthropic/claude-opus-4-5") is True
|
||||
assert should_use_oauth_provider("sk-ant-oat01-xxx", "anthropic/claude-opus-4-7") is True
|
||||
assert should_use_oauth_provider("sk-ant-oat01-xxx", "claude-sonnet-4") is True
|
||||
|
||||
|
||||
def test_should_not_use_oauth_for_regular_key():
|
||||
"""Regular API keys should not use OAuth provider."""
|
||||
assert should_use_oauth_provider("sk-ant-api03-xxx", "claude-opus-4-5") is False
|
||||
assert should_use_oauth_provider("sk-or-v1-xxx", "anthropic/claude-opus-4-5") is False
|
||||
assert should_use_oauth_provider("sk-ant-api03-xxx", "claude-opus-4-7") is False
|
||||
assert should_use_oauth_provider("sk-or-v1-xxx", "anthropic/claude-opus-4-7") is False
|
||||
|
||||
|
||||
def test_should_not_use_oauth_for_non_anthropic():
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
"""Test SessionManager audit log functionality."""
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.session.manager import Session, SessionManager
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def session_manager(tmp_path):
|
||||
return SessionManager(workspace=tmp_path)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def session():
|
||||
s = Session(key="telegram:12345")
|
||||
s.add_message("user", "Hello")
|
||||
s.add_message("assistant", "Hi there!")
|
||||
return s
|
||||
|
||||
|
||||
def test_save_creates_audit_file(session_manager, session):
|
||||
"""SessionManager.save() should create a monthly audit log file."""
|
||||
session_manager.save(session)
|
||||
|
||||
audit_files = list(session_manager.sessions_dir.glob("*.audit.*.jsonl"))
|
||||
assert len(audit_files) == 1
|
||||
assert "telegram_12345.audit." in audit_files[0].name
|
||||
|
||||
|
||||
def test_audit_file_contains_save_marker(session_manager, session):
|
||||
"""Audit log should start with a save_marker line containing metadata."""
|
||||
session_manager.save(session)
|
||||
|
||||
audit_files = list(session_manager.sessions_dir.glob("*.audit.*.jsonl"))
|
||||
lines = audit_files[0].read_text().strip().split("\n")
|
||||
|
||||
marker = json.loads(lines[0])
|
||||
assert marker["_type"] == "save_marker"
|
||||
assert marker["message_count"] == 2
|
||||
assert "timestamp" in marker
|
||||
|
||||
|
||||
def test_audit_file_contains_all_messages(session_manager, session):
|
||||
"""Audit log should contain all session messages after the save marker."""
|
||||
session_manager.save(session)
|
||||
|
||||
audit_files = list(session_manager.sessions_dir.glob("*.audit.*.jsonl"))
|
||||
lines = audit_files[0].read_text().strip().split("\n")
|
||||
|
||||
# Line 0 = save_marker, lines 1-2 = messages
|
||||
assert len(lines) == 3
|
||||
msg1 = json.loads(lines[1])
|
||||
msg2 = json.loads(lines[2])
|
||||
assert msg1["role"] == "user"
|
||||
assert msg1["content"] == "Hello"
|
||||
assert msg2["role"] == "assistant"
|
||||
assert msg2["content"] == "Hi there!"
|
||||
|
||||
|
||||
def test_audit_file_is_append_only(session_manager, session):
|
||||
"""Multiple saves should append to the same audit file, not overwrite."""
|
||||
session_manager.save(session)
|
||||
|
||||
# Add another message and save again
|
||||
session.add_message("user", "How are you?")
|
||||
session_manager.save(session)
|
||||
|
||||
audit_files = list(session_manager.sessions_dir.glob("*.audit.*.jsonl"))
|
||||
assert len(audit_files) == 1 # Same file
|
||||
|
||||
lines = audit_files[0].read_text().strip().split("\n")
|
||||
|
||||
# First save: 1 marker + 2 messages = 3 lines
|
||||
# Second save: 1 marker + 3 messages = 4 lines
|
||||
# Total: 7 lines
|
||||
assert len(lines) == 7
|
||||
|
||||
# Both save markers present
|
||||
markers = [json.loads(l) for l in lines if json.loads(l).get("_type") == "save_marker"]
|
||||
assert len(markers) == 2
|
||||
assert markers[0]["message_count"] == 2
|
||||
assert markers[1]["message_count"] == 3
|
||||
|
||||
|
||||
def test_audit_preserves_message_fields(session_manager):
|
||||
"""Audit log should preserve all message fields including reasoning_content."""
|
||||
session = Session(key="test:preserve")
|
||||
session.add_raw_message({
|
||||
"role": "assistant",
|
||||
"content": "thinking response",
|
||||
"reasoning_content": [{"type": "thinking", "thinking": "deep thoughts"}],
|
||||
"timestamp": "2026-03-22T12:00:00",
|
||||
})
|
||||
|
||||
session_manager.save(session)
|
||||
|
||||
audit_files = list(session_manager.sessions_dir.glob("*.audit.*.jsonl"))
|
||||
lines = audit_files[0].read_text().strip().split("\n")
|
||||
|
||||
msg = json.loads(lines[1])
|
||||
assert msg["reasoning_content"] == [{"type": "thinking", "thinking": "deep thoughts"}]
|
||||
|
||||
|
||||
def test_audit_failure_does_not_break_save(session_manager, session, tmp_path):
|
||||
"""If audit logging fails, the main session save should still succeed.
|
||||
|
||||
_append_audit has its own try/except, so internal failures are caught.
|
||||
We simulate a realistic failure by making the sessions dir read-only
|
||||
for audit file creation.
|
||||
"""
|
||||
# First save works (creates both session file and audit file)
|
||||
session_manager.save(session)
|
||||
|
||||
path = session_manager._get_session_path(session.key)
|
||||
assert path.exists()
|
||||
|
||||
# Remove audit files and make a blocking file at the audit path
|
||||
# so the next audit open("a") fails
|
||||
for af in session_manager.sessions_dir.glob("*.audit.*.jsonl"):
|
||||
af.unlink()
|
||||
|
||||
# Create a directory where the audit file should be — open() will fail
|
||||
from datetime import datetime
|
||||
now = datetime.now()
|
||||
bad_path = session_manager.sessions_dir / f"telegram_12345.audit.{now:%Y-%m}.jsonl"
|
||||
bad_path.mkdir()
|
||||
|
||||
# Second save should succeed despite audit failure
|
||||
session.add_message("user", "another message")
|
||||
session_manager.save(session)
|
||||
|
||||
# Session file should still be written correctly
|
||||
with open(path) as f:
|
||||
first_line = json.loads(f.readline())
|
||||
assert first_line["_type"] == "metadata"
|
||||
@@ -0,0 +1,139 @@
|
||||
# tests/test_subagent_wait.py
|
||||
"""Tests for wait_for_subagents with top-level and child subagents."""
|
||||
import pytest
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wait_for_top_level_subagent():
|
||||
"""Test that wait_for works for top-level subagents spawned from telegram."""
|
||||
bus = MessageBus()
|
||||
provider = MagicMock(spec=LLMProvider)
|
||||
provider.chat = AsyncMock(return_value=LLMResponse(
|
||||
content="Task completed",
|
||||
tool_calls=[]
|
||||
))
|
||||
provider.get_default_model = MagicMock(return_value="test-model")
|
||||
provider.thinking_budget = 0
|
||||
|
||||
workspace = Path("/tmp/test-subagent")
|
||||
workspace.mkdir(exist_ok=True)
|
||||
|
||||
manager = SubagentManager(
|
||||
bus=bus,
|
||||
provider=provider,
|
||||
workspace=workspace
|
||||
)
|
||||
|
||||
# Spawn a top-level subagent (origin channel = "telegram")
|
||||
task_id = await manager.spawn(
|
||||
task="Test task",
|
||||
label="test",
|
||||
model=None,
|
||||
origin_channel="telegram",
|
||||
origin_chat_id="12345"
|
||||
)
|
||||
|
||||
# Wait for it to complete
|
||||
result = await manager.wait_for([task_id])
|
||||
|
||||
# Should find the result (not "No result found")
|
||||
assert "No result found" not in result
|
||||
assert task_id in result
|
||||
assert "Task completed" in result or "completed" in result.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wait_for_child_subagent():
|
||||
"""Test that wait_for works for child subagents (orchestrator pattern)."""
|
||||
bus = MessageBus()
|
||||
provider = MagicMock(spec=LLMProvider)
|
||||
provider.chat = AsyncMock(return_value=LLMResponse(
|
||||
content="Child task completed",
|
||||
tool_calls=[]
|
||||
))
|
||||
provider.get_default_model = MagicMock(return_value="test-model")
|
||||
provider.thinking_budget = 0
|
||||
|
||||
workspace = Path("/tmp/test-subagent")
|
||||
workspace.mkdir(exist_ok=True)
|
||||
|
||||
manager = SubagentManager(
|
||||
bus=bus,
|
||||
provider=provider,
|
||||
workspace=workspace
|
||||
)
|
||||
|
||||
# Spawn a child subagent (origin channel = "subagent")
|
||||
task_id = await manager.spawn(
|
||||
task="Child test task",
|
||||
label="test-child",
|
||||
model=None,
|
||||
origin_channel="subagent",
|
||||
origin_chat_id="parent-id"
|
||||
)
|
||||
|
||||
# Wait for it to complete
|
||||
result = await manager.wait_for([task_id])
|
||||
|
||||
# Should find the result (not "No result found")
|
||||
assert "No result found" not in result
|
||||
assert task_id in result
|
||||
assert "Child task completed" in result or "completed" in result.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wait_for_multiple_subagents():
|
||||
"""Test waiting for multiple subagents of different types."""
|
||||
bus = MessageBus()
|
||||
call_count = 0
|
||||
|
||||
async def chat_response(*args, **kwargs):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return LLMResponse(content=f"Task {call_count} completed", tool_calls=[])
|
||||
|
||||
provider = MagicMock(spec=LLMProvider)
|
||||
provider.chat = AsyncMock(side_effect=chat_response)
|
||||
provider.get_default_model = MagicMock(return_value="test-model")
|
||||
provider.thinking_budget = 0
|
||||
|
||||
workspace = Path("/tmp/test-subagent")
|
||||
workspace.mkdir(exist_ok=True)
|
||||
|
||||
manager = SubagentManager(
|
||||
bus=bus,
|
||||
provider=provider,
|
||||
workspace=workspace
|
||||
)
|
||||
|
||||
# Spawn one top-level and one child subagent
|
||||
task_id_1 = await manager.spawn(
|
||||
task="Top-level task",
|
||||
label="test-top",
|
||||
model=None,
|
||||
origin_channel="telegram",
|
||||
origin_chat_id="12345"
|
||||
)
|
||||
|
||||
task_id_2 = await manager.spawn(
|
||||
task="Child task",
|
||||
label="test-child",
|
||||
model=None,
|
||||
origin_channel="subagent",
|
||||
origin_chat_id="parent"
|
||||
)
|
||||
|
||||
# Wait for both
|
||||
result = await manager.wait_for([task_id_1, task_id_2])
|
||||
|
||||
# Should find both results
|
||||
assert "No result found" not in result
|
||||
assert task_id_1 in result
|
||||
assert task_id_2 in result
|
||||
assert "Task 1 completed" in result or "completed" in result.lower()
|
||||
assert "Task 2 completed" in result or "completed" in result.lower()
|
||||
@@ -1,167 +0,0 @@
|
||||
"""Tests for /stop task cancellation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _make_loop():
|
||||
"""Create a minimal AgentLoop with mocked dependencies."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.queue import MessageBus
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
workspace = MagicMock()
|
||||
workspace.__truediv__ = MagicMock(return_value=MagicMock())
|
||||
|
||||
with patch("nanobot.agent.loop.ContextBuilder"), \
|
||||
patch("nanobot.agent.loop.SessionManager"), \
|
||||
patch("nanobot.agent.loop.SubagentManager") as MockSubMgr:
|
||||
MockSubMgr.return_value.cancel_by_session = AsyncMock(return_value=0)
|
||||
loop = AgentLoop(bus=bus, provider=provider, workspace=workspace)
|
||||
return loop, bus
|
||||
|
||||
|
||||
class TestHandleStop:
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_no_active_task(self):
|
||||
from nanobot.bus.events import InboundMessage
|
||||
|
||||
loop, bus = _make_loop()
|
||||
msg = InboundMessage(channel="test", sender_id="u1", chat_id="c1", content="/stop")
|
||||
await loop._handle_stop(msg)
|
||||
out = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||
assert "No active task" in out.content
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_cancels_active_task(self):
|
||||
from nanobot.bus.events import InboundMessage
|
||||
|
||||
loop, bus = _make_loop()
|
||||
cancelled = asyncio.Event()
|
||||
|
||||
async def slow_task():
|
||||
try:
|
||||
await asyncio.sleep(60)
|
||||
except asyncio.CancelledError:
|
||||
cancelled.set()
|
||||
raise
|
||||
|
||||
task = asyncio.create_task(slow_task())
|
||||
await asyncio.sleep(0)
|
||||
loop._active_tasks["test:c1"] = [task]
|
||||
|
||||
msg = InboundMessage(channel="test", sender_id="u1", chat_id="c1", content="/stop")
|
||||
await loop._handle_stop(msg)
|
||||
|
||||
assert cancelled.is_set()
|
||||
out = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||
assert "stopped" in out.content.lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_cancels_multiple_tasks(self):
|
||||
from nanobot.bus.events import InboundMessage
|
||||
|
||||
loop, bus = _make_loop()
|
||||
events = [asyncio.Event(), asyncio.Event()]
|
||||
|
||||
async def slow(idx):
|
||||
try:
|
||||
await asyncio.sleep(60)
|
||||
except asyncio.CancelledError:
|
||||
events[idx].set()
|
||||
raise
|
||||
|
||||
tasks = [asyncio.create_task(slow(i)) for i in range(2)]
|
||||
await asyncio.sleep(0)
|
||||
loop._active_tasks["test:c1"] = tasks
|
||||
|
||||
msg = InboundMessage(channel="test", sender_id="u1", chat_id="c1", content="/stop")
|
||||
await loop._handle_stop(msg)
|
||||
|
||||
assert all(e.is_set() for e in events)
|
||||
out = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||
assert "2 task" in out.content
|
||||
|
||||
|
||||
class TestDispatch:
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_processes_and_publishes(self):
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
|
||||
loop, bus = _make_loop()
|
||||
msg = InboundMessage(channel="test", sender_id="u1", chat_id="c1", content="hello")
|
||||
loop._process_message = AsyncMock(
|
||||
return_value=OutboundMessage(channel="test", chat_id="c1", content="hi")
|
||||
)
|
||||
await loop._dispatch(msg)
|
||||
out = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||
assert out.content == "hi"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_processing_lock_serializes(self):
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
|
||||
loop, bus = _make_loop()
|
||||
order = []
|
||||
|
||||
async def mock_process(m, **kwargs):
|
||||
order.append(f"start-{m.content}")
|
||||
await asyncio.sleep(0.05)
|
||||
order.append(f"end-{m.content}")
|
||||
return OutboundMessage(channel="test", chat_id="c1", content=m.content)
|
||||
|
||||
loop._process_message = mock_process
|
||||
msg1 = InboundMessage(channel="test", sender_id="u1", chat_id="c1", content="a")
|
||||
msg2 = InboundMessage(channel="test", sender_id="u1", chat_id="c1", content="b")
|
||||
|
||||
t1 = asyncio.create_task(loop._dispatch(msg1))
|
||||
t2 = asyncio.create_task(loop._dispatch(msg2))
|
||||
await asyncio.gather(t1, t2)
|
||||
assert order == ["start-a", "end-a", "start-b", "end-b"]
|
||||
|
||||
|
||||
class TestSubagentCancellation:
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_by_session(self):
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
from nanobot.bus.queue import MessageBus
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(provider=provider, workspace=MagicMock(), bus=bus)
|
||||
|
||||
cancelled = asyncio.Event()
|
||||
|
||||
async def slow():
|
||||
try:
|
||||
await asyncio.sleep(60)
|
||||
except asyncio.CancelledError:
|
||||
cancelled.set()
|
||||
raise
|
||||
|
||||
task = asyncio.create_task(slow())
|
||||
await asyncio.sleep(0)
|
||||
mgr._running_tasks["sub-1"] = task
|
||||
mgr._session_tasks["test:c1"] = {"sub-1"}
|
||||
|
||||
count = await mgr.cancel_by_session("test:c1")
|
||||
assert count == 1
|
||||
assert cancelled.is_set()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_by_session_no_tasks(self):
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
from nanobot.bus.queue import MessageBus
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
mgr = SubagentManager(provider=provider, workspace=MagicMock(), bus=bus)
|
||||
assert await mgr.cancel_by_session("nonexistent") == 0
|
||||
Reference in New Issue
Block a user