Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2523db4da7 | ||
|
|
470b7448ac | ||
|
|
b5b2346db5 | ||
|
|
617ff61347 | ||
|
|
8609181447 |
@@ -184,18 +184,28 @@ cat ../webber-sandbox/TASKS.md
|
|||||||
## Versioning & Releases
|
## Versioning & Releases
|
||||||
|
|
||||||
Uses prefixed tags:
|
Uses prefixed tags:
|
||||||
- `api/v0.3.0` → Triggers API Docker build
|
- `api/vX.Y.Z` → Triggers API Docker build
|
||||||
- `cli/v0.1.0` → Triggers CLI build (future)
|
- `cli/vX.Y.Z` → Triggers CLI build (future)
|
||||||
|
|
||||||
|
### MANDATORY Release Procedure
|
||||||
|
|
||||||
|
**NEVER push a tag before updating version files.** Follow this exact order:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# API release
|
# 1. Update version in pyproject.toml
|
||||||
cd webber-api
|
# 2. Update CHANGELOG.md with release notes
|
||||||
# Update version in pyproject.toml
|
# 3. Commit the version bump
|
||||||
git add -A && git commit -m "chore: release api v0.3.0"
|
git add -A && git commit -m "chore: release api vX.Y.Z"
|
||||||
git tag api/v0.3.0
|
|
||||||
|
# 4. Create the tag (AFTER the commit)
|
||||||
|
git tag api/vX.Y.Z
|
||||||
|
|
||||||
|
# 5. Push everything together
|
||||||
git push origin main --tags
|
git push origin main --tags
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**Why this matters:** Pushing a tag before the version commit requires deleting and recreating the tag, which can trigger CI/CD pipelines prematurely and cause deployment issues.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Troubleshooting
|
## Troubleshooting
|
||||||
|
|||||||
@@ -7,6 +7,58 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
|
## [0.4.0] - 2026-01-11
|
||||||
|
|
||||||
|
### Added
|
||||||
|
- Conversation persistence layer with SQLAlchemy async
|
||||||
|
- Database models: `Conversation`, `Message` with UUID primary keys
|
||||||
|
- SQLite (dev) and PostgreSQL (prod) support via async engines
|
||||||
|
- Lazy database initialization pattern
|
||||||
|
- Context management infrastructure
|
||||||
|
- Token counting utilities using `litellm`
|
||||||
|
- Context summarization at 80% token threshold
|
||||||
|
- XML-tagged context prompt building for agent injection
|
||||||
|
- REST API for multi-turn conversations
|
||||||
|
- `POST /conversations/` - Create new conversation
|
||||||
|
- `GET /conversations/` - List conversations
|
||||||
|
- `GET /conversations/{id}` - Get conversation with history
|
||||||
|
- `POST /conversations/{id}/messages` - Add message (triggers agent)
|
||||||
|
- `DELETE /conversations/{id}` - Delete conversation
|
||||||
|
- New dependencies: `sqlalchemy[asyncio]~=2.0.36`, `aiosqlite~=0.21.0`, `litellm~=1.57.0`
|
||||||
|
- Config settings: `database_url`, `summarization_threshold`, `keep_recent_messages`
|
||||||
|
- 19 conversation tests, 6 token counting tests (176 total tests passing)
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
- Updated COVERAGE.md to ~80% complete
|
||||||
|
- Quieter pytest output (`-q --tb=short` instead of `-v`)
|
||||||
|
|
||||||
|
## [0.3.4] - 2026-01-11
|
||||||
|
|
||||||
|
### Added
|
||||||
|
- Task Agent - Full orchestrator for autonomous multi-step task execution
|
||||||
|
- Has ALL tools: read, write, edit, bash (full), web_search
|
||||||
|
- New `spawn_agent` tool to launch sub-agents (Explore, Plan) for focused work
|
||||||
|
- Recursion prevention: cannot spawn nested Task agents
|
||||||
|
- 22 unit tests for registration, tools, spawn_agent, and API
|
||||||
|
- Complete agent hierarchy: Explore (read-only) → Plan (read-only) → Task (orchestrator)
|
||||||
|
|
||||||
|
## [0.3.3] - 2026-01-11
|
||||||
|
|
||||||
|
### Added
|
||||||
|
- Plan Agent - READ-ONLY software architect that designs implementation strategies
|
||||||
|
- Uses only read-only tools: `read_file`, `glob_files`, `grep_content`, `bash_readonly`
|
||||||
|
- Creates step-by-step implementation plans with critical files list
|
||||||
|
- 15 unit tests for registration, tools, and API
|
||||||
|
- Web search summarizer added to roadmap (future feature)
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
- Updated COVERAGE.md to ~70% complete
|
||||||
|
|
||||||
|
## [0.3.2] - 2026-01-11
|
||||||
|
|
||||||
|
### Added
|
||||||
|
- Mandatory release procedure documentation in AGENTS.md
|
||||||
|
|
||||||
## [0.3.1] - 2026-01-11
|
## [0.3.1] - 2026-01-11
|
||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|||||||
+55
-11
@@ -2,7 +2,7 @@
|
|||||||
|
|
||||||
> Tracking progress towards Claude Code-like functionality
|
> Tracking progress towards Claude Code-like functionality
|
||||||
|
|
||||||
## Current Status: ~65% Complete
|
## Current Status: ~80% Complete
|
||||||
|
|
||||||
Last updated: 2026-01-11
|
Last updated: 2026-01-11
|
||||||
|
|
||||||
@@ -44,6 +44,20 @@ Last updated: 2026-01-11
|
|||||||
|
|
||||||
**Gap:** Mistral Nemo sometimes hallucinates instead of using tool results.
|
**Gap:** Mistral Nemo sometimes hallucinates instead of using tool results.
|
||||||
|
|
||||||
|
### Phase 2b: Plan Agent ✅ Complete
|
||||||
|
|
||||||
|
| Component | Status | Notes |
|
||||||
|
|-----------|--------|-------|
|
||||||
|
| `PlanAgentImpl` | ✅ | READ-ONLY software architect agent |
|
||||||
|
| System prompts | ✅ | Architecture-focused with tool examples |
|
||||||
|
| Tool registration | ✅ | Only read-only tools (4 tools) |
|
||||||
|
| Streaming support | ✅ | `run_stream()` method with SSE |
|
||||||
|
| Unit tests | ✅ | 15 tests for registration, tools, API |
|
||||||
|
|
||||||
|
**Available tools:** `read_file`, `glob_files`, `grep_content`, `bash_readonly` (read-only only)
|
||||||
|
|
||||||
|
**Purpose:** Design implementation strategies before coding - explores codebase and creates step-by-step plans.
|
||||||
|
|
||||||
### Phase 3: CLI Foundation ✅ Complete
|
### Phase 3: CLI Foundation ✅ Complete
|
||||||
|
|
||||||
| Component | Status | Notes |
|
| Component | Status | Notes |
|
||||||
@@ -54,16 +68,19 @@ Last updated: 2026-01-11
|
|||||||
| Markdown rendering | ✅ | Rich markdown output |
|
| Markdown rendering | ✅ | Rich markdown output |
|
||||||
| Streaming display | ✅ | Real-time token output with `--stream` flag |
|
| Streaming display | ✅ | Real-time token output with `--stream` flag |
|
||||||
|
|
||||||
### Phase 4: Agentic Loop ⚠️ Partial
|
### Phase 4: Agentic Loop ✅ Complete
|
||||||
|
|
||||||
| Component | Status | Notes |
|
| Component | Status | Notes |
|
||||||
|-----------|--------|-------|
|
|-----------|--------|-------|
|
||||||
| `webber-cli chat` command | ✅ | Interactive mode with streaming |
|
| `webber-cli chat` command | ✅ | Interactive mode with streaming |
|
||||||
| `webber-cli explore` command | ✅ | One-shot query with streaming |
|
| `webber-cli explore` command | ✅ | One-shot query with streaming |
|
||||||
| `SessionState` dataclass | ✅ | Basic context tracking |
|
| `SessionState` dataclass | ✅ | Basic context tracking |
|
||||||
| `AgenticLoop` class | ⚠️ | Basic implementation, not fully utilized |
|
| `AgenticLoop` class | ✅ | Basic implementation |
|
||||||
| Conversation history | ❌ | Not persisted between turns in CLI |
|
| Conversation persistence | ✅ | SQLAlchemy async with SQLite/PostgreSQL |
|
||||||
| Context management | ❌ | No token counting or summarization |
|
| Context summarization | ✅ | Token counting (litellm) + auto-summarization |
|
||||||
|
| Conversation API | ✅ | `/conversations/` REST endpoints |
|
||||||
|
|
||||||
|
**Database:** SQLite (dev) or PostgreSQL (prod), async via SQLAlchemy 2.0
|
||||||
|
|
||||||
### Phase 5: REST API ✅ Complete
|
### Phase 5: REST API ✅ Complete
|
||||||
|
|
||||||
@@ -94,15 +111,16 @@ Last updated: 2026-01-11
|
|||||||
|
|
||||||
| Feature | Category | Description | Complexity |
|
| Feature | Category | Description | Complexity |
|
||||||
|---------|----------|-------------|------------|
|
|---------|----------|-------------|------------|
|
||||||
| **Plan Agent** | Agents | Design implementation approaches | High |
|
| ~~**Plan Agent**~~ | Agents | ✅ Design implementation approaches | High |
|
||||||
| **Task Agent** | Agents | Autonomous multi-step execution | High |
|
| ~~**Task Agent**~~ | Agents | ✅ Autonomous multi-step execution | High |
|
||||||
| **Context summarization** | Infrastructure | Compress history at token limit | High |
|
| ~~**Context summarization**~~ | Infrastructure | ✅ Token counting + auto-summarization | High |
|
||||||
| **Conversation persistence** | CLI | Multi-turn memory in chat mode | Medium |
|
| ~~**Conversation persistence**~~ | Infrastructure | ✅ SQLAlchemy async database layer | Medium |
|
||||||
|
|
||||||
### Medium Priority
|
### Medium Priority
|
||||||
|
|
||||||
| Feature | Category | Description | Complexity |
|
| Feature | Category | Description | Complexity |
|
||||||
|---------|----------|-------------|------------|
|
|---------|----------|-------------|------------|
|
||||||
|
| **Web search summarizer** | Tools | Agent to extract core content from web pages (remove nav, footers, etc.) and preserve relevant links for nested fetching | Medium |
|
||||||
| **Tool result caching** | Infrastructure | Cache file reads for performance | Low |
|
| **Tool result caching** | Infrastructure | Cache file reads for performance | Low |
|
||||||
| **Session persistence** | CLI | Save/resume conversations | Medium |
|
| **Session persistence** | CLI | Save/resume conversations | Medium |
|
||||||
| **Todo tracking** | CLI | Built-in task list (`/todo`) | Medium |
|
| **Todo tracking** | CLI | Built-in task list (`/todo`) | Medium |
|
||||||
@@ -129,10 +147,16 @@ Last updated: 2026-01-11
|
|||||||
|------|---------|--------|--------|
|
|------|---------|--------|--------|
|
||||||
| Tool unit tests | 109 | 109 | ✅ |
|
| Tool unit tests | 109 | 109 | ✅ |
|
||||||
| API tests | 11 | 11 | ✅ |
|
| API tests | 11 | 11 | ✅ |
|
||||||
|
| Plan agent tests | 15 | 15 | ✅ |
|
||||||
|
| Task agent tests | 15 | 15 | ✅ |
|
||||||
|
| Conversation tests | 19 | 19 | ✅ |
|
||||||
|
| Token tests | 6 | 6 | ✅ |
|
||||||
| Security tests | 14 | 14 | ✅ |
|
| Security tests | 14 | 14 | ✅ |
|
||||||
| Integration tests | 10 | 10 | ✅ Agent + real LLM |
|
| Integration tests | 10 | 10 | ✅ Agent + real LLM |
|
||||||
| E2E tests | 12 | 12 | ✅ Full API workflow |
|
| E2E tests | 12 | 12 | ✅ Full API workflow |
|
||||||
|
|
||||||
|
**Total: 176 tests passing**
|
||||||
|
|
||||||
**Test breakdown:**
|
**Test breakdown:**
|
||||||
- Read/Glob/Grep tools: 17 tests
|
- Read/Glob/Grep tools: 17 tests
|
||||||
- Edit/Write tools: 22 tests
|
- Edit/Write tools: 22 tests
|
||||||
@@ -140,6 +164,10 @@ Last updated: 2026-01-11
|
|||||||
- Web search: 10 tests
|
- Web search: 10 tests
|
||||||
- Gitignore filtering: 10 tests
|
- Gitignore filtering: 10 tests
|
||||||
- API endpoints: 11 tests
|
- API endpoints: 11 tests
|
||||||
|
- Plan agent: 15 tests
|
||||||
|
- Task agent: 15 tests
|
||||||
|
- Conversations: 19 tests
|
||||||
|
- Tokens: 6 tests
|
||||||
- Security: 14 tests
|
- Security: 14 tests
|
||||||
- Health checks: 2 tests
|
- Health checks: 2 tests
|
||||||
- Integration (LLM): 10 tests
|
- Integration (LLM): 10 tests
|
||||||
@@ -166,9 +194,9 @@ pytest tests/ --run-integration --run-e2e
|
|||||||
|
|
||||||
1. **Model hallucination** - Mistral Nemo sometimes makes up file contents instead of using actual tool results.
|
1. **Model hallucination** - Mistral Nemo sometimes makes up file contents instead of using actual tool results.
|
||||||
|
|
||||||
2. **No conversation memory** - CLI chat mode doesn't persist context between sessions.
|
2. **Temperature setting** - Changed from 0.0 to 0.3 for Mistral Nemo compatibility, may affect determinism.
|
||||||
|
|
||||||
3. **Temperature setting** - Changed from 0.0 to 0.3 for Mistral Nemo compatibility, may affect determinism.
|
3. **SQLAlchemy deprecation** - `datetime.utcnow()` deprecation warning from SQLAlchemy.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -205,10 +233,26 @@ curl -X POST http://localhost:8095/agents/run \
|
|||||||
-H "Content-Type: application/json" \
|
-H "Content-Type: application/json" \
|
||||||
-d '{"agent_type":"explore","prompt":"list python files","working_dir":"."}'
|
-d '{"agent_type":"explore","prompt":"list python files","working_dir":"."}'
|
||||||
|
|
||||||
|
# Plan agent (read-only, creates implementation plans)
|
||||||
|
curl -X POST http://localhost:8095/agents/run \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"agent_type":"plan","prompt":"plan how to add user auth","working_dir":"."}'
|
||||||
|
|
||||||
# Streaming endpoint
|
# Streaming endpoint
|
||||||
curl -N http://localhost:8095/agents/stream \
|
curl -N http://localhost:8095/agents/stream \
|
||||||
-H "Content-Type: application/json" \
|
-H "Content-Type: application/json" \
|
||||||
-d '{"agent_type":"explore","prompt":"find config files","working_dir":"."}'
|
-d '{"agent_type":"explore","prompt":"find config files","working_dir":"."}'
|
||||||
|
|
||||||
|
# Conversation API (stateful multi-turn)
|
||||||
|
curl -X POST http://localhost:8095/conversations/ \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-H "X-API-Key: dev-key" \
|
||||||
|
-d '{"agent_type":"explore","working_dir":"."}'
|
||||||
|
|
||||||
|
curl -X POST http://localhost:8095/conversations/{id}/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-H "X-API-Key: dev-key" \
|
||||||
|
-d '{"content":"find all Python files"}'
|
||||||
```
|
```
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "webber-api"
|
name = "webber-api"
|
||||||
version = "0.3.1"
|
version = "0.4.0"
|
||||||
description = "Webber API - Multi-Agent AI Development Server"
|
description = "Webber API - Multi-Agent AI Development Server"
|
||||||
authors = [
|
authors = [
|
||||||
{name = "jpmschweitzer"}
|
{name = "jpmschweitzer"}
|
||||||
@@ -27,7 +27,7 @@ include = ["src*"]
|
|||||||
testpaths = ["tests"]
|
testpaths = ["tests"]
|
||||||
python_files = ["test_*.py"]
|
python_files = ["test_*.py"]
|
||||||
python_functions = ["test_*"]
|
python_functions = ["test_*"]
|
||||||
addopts = "-v --strict-markers"
|
addopts = "-q --strict-markers --tb=short"
|
||||||
markers = [
|
markers = [
|
||||||
"integration: marks tests as integration tests (require Ollama to be running)",
|
"integration: marks tests as integration tests (require Ollama to be running)",
|
||||||
"e2e: marks tests as end-to-end tests (require API server to be running)",
|
"e2e: marks tests as end-to-end tests (require API server to be running)",
|
||||||
|
|||||||
@@ -25,3 +25,10 @@ rich~=13.9.0
|
|||||||
python-multipart~=0.0.21
|
python-multipart~=0.0.21
|
||||||
python-dotenv~=1.2.1
|
python-dotenv~=1.2.1
|
||||||
pathspec~=0.12.1 # Gitignore pattern matching
|
pathspec~=0.12.1 # Gitignore pattern matching
|
||||||
|
|
||||||
|
# Database
|
||||||
|
sqlalchemy[asyncio]~=2.0.36
|
||||||
|
aiosqlite~=0.21.0 # SQLite async driver (dev)
|
||||||
|
|
||||||
|
# Token counting
|
||||||
|
litellm~=1.57.0 # Multi-model token counting
|
||||||
|
|||||||
@@ -0,0 +1,14 @@
|
|||||||
|
"""
|
||||||
|
Database package for Webber.
|
||||||
|
|
||||||
|
Provides async SQLAlchemy database access following core-api patterns.
|
||||||
|
"""
|
||||||
|
from src.db.database import Database, get_database, get_session
|
||||||
|
from src.db.models import Base
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Database",
|
||||||
|
"get_database",
|
||||||
|
"get_session",
|
||||||
|
"Base",
|
||||||
|
]
|
||||||
@@ -0,0 +1,123 @@
|
|||||||
|
"""
|
||||||
|
Async SQLAlchemy database management.
|
||||||
|
|
||||||
|
Pattern from core-api: singleton Database class with async session factory.
|
||||||
|
"""
|
||||||
|
from collections.abc import AsyncGenerator
|
||||||
|
from functools import lru_cache
|
||||||
|
|
||||||
|
from sqlalchemy.ext.asyncio import (
|
||||||
|
AsyncEngine,
|
||||||
|
AsyncSession,
|
||||||
|
async_sessionmaker,
|
||||||
|
create_async_engine,
|
||||||
|
)
|
||||||
|
|
||||||
|
from src.shared.config import get_settings
|
||||||
|
from src.shared.logging import get_logger
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class Database:
|
||||||
|
"""
|
||||||
|
Async database connection manager.
|
||||||
|
|
||||||
|
Manages SQLAlchemy async engine and session factory.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, url: str):
|
||||||
|
"""
|
||||||
|
Initialize database with connection URL.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
url: SQLAlchemy async connection URL
|
||||||
|
e.g., "sqlite+aiosqlite:///./webber.db"
|
||||||
|
or "postgresql+asyncpg://user:pass@host/db"
|
||||||
|
"""
|
||||||
|
self._url = url
|
||||||
|
self._engine: AsyncEngine | None = None
|
||||||
|
self._session_factory: async_sessionmaker[AsyncSession] | None = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def engine(self) -> AsyncEngine:
|
||||||
|
"""Get or create the async engine."""
|
||||||
|
if self._engine is None:
|
||||||
|
self._engine = create_async_engine(
|
||||||
|
self._url,
|
||||||
|
echo=get_settings().debug,
|
||||||
|
pool_pre_ping=True,
|
||||||
|
)
|
||||||
|
return self._engine
|
||||||
|
|
||||||
|
@property
|
||||||
|
def session_factory(self) -> async_sessionmaker[AsyncSession]:
|
||||||
|
"""Get or create the session factory."""
|
||||||
|
if self._session_factory is None:
|
||||||
|
self._session_factory = async_sessionmaker(
|
||||||
|
bind=self.engine,
|
||||||
|
class_=AsyncSession,
|
||||||
|
expire_on_commit=False,
|
||||||
|
autoflush=False,
|
||||||
|
)
|
||||||
|
return self._session_factory
|
||||||
|
|
||||||
|
async def create_tables(self) -> None:
|
||||||
|
"""Create all tables (for development)."""
|
||||||
|
from src.db.models import Base
|
||||||
|
|
||||||
|
async with self.engine.begin() as conn:
|
||||||
|
await conn.run_sync(Base.metadata.create_all)
|
||||||
|
logger.info("Database tables created")
|
||||||
|
|
||||||
|
async def close(self) -> None:
|
||||||
|
"""Close the database connection."""
|
||||||
|
if self._engine:
|
||||||
|
await self._engine.dispose()
|
||||||
|
self._engine = None
|
||||||
|
self._session_factory = None
|
||||||
|
logger.info("Database connection closed")
|
||||||
|
|
||||||
|
|
||||||
|
# Singleton instance
|
||||||
|
_database: Database | None = None
|
||||||
|
_tables_created: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache
|
||||||
|
def get_database() -> Database:
|
||||||
|
"""Get the database singleton."""
|
||||||
|
global _database
|
||||||
|
if _database is None:
|
||||||
|
settings = get_settings()
|
||||||
|
_database = Database(settings.database_url)
|
||||||
|
return _database
|
||||||
|
|
||||||
|
|
||||||
|
async def _ensure_tables() -> None:
|
||||||
|
"""Ensure database tables exist (lazy initialization)."""
|
||||||
|
global _tables_created
|
||||||
|
if not _tables_created:
|
||||||
|
database = get_database()
|
||||||
|
await database.create_tables()
|
||||||
|
_tables_created = True
|
||||||
|
|
||||||
|
|
||||||
|
async def get_session() -> AsyncGenerator[AsyncSession, None]:
|
||||||
|
"""
|
||||||
|
Dependency for getting async database sessions.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@router.get("/")
|
||||||
|
async def endpoint(session: AsyncSession = Depends(get_session)):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
await _ensure_tables()
|
||||||
|
database = get_database()
|
||||||
|
async with database.session_factory() as session:
|
||||||
|
try:
|
||||||
|
yield session
|
||||||
|
await session.commit()
|
||||||
|
except Exception:
|
||||||
|
await session.rollback()
|
||||||
|
raise
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
"""
|
||||||
|
SQLAlchemy Base model for all database models.
|
||||||
|
"""
|
||||||
|
from sqlalchemy.orm import DeclarativeBase
|
||||||
|
|
||||||
|
|
||||||
|
class Base(DeclarativeBase):
|
||||||
|
"""Base class for all SQLAlchemy models."""
|
||||||
|
pass
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
"""
|
||||||
|
Plan Agent - Software architect for implementation planning.
|
||||||
|
|
||||||
|
The Plan agent explores codebases and designs step-by-step implementation
|
||||||
|
strategies. It uses only read-only tools and cannot modify any files.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
from src.domains.agents.plan import plan_agent, plan
|
||||||
|
|
||||||
|
# Direct agent access
|
||||||
|
result = await plan_agent.run("Plan how to add user authentication")
|
||||||
|
|
||||||
|
# Convenience function
|
||||||
|
result = await plan("Plan how to add user authentication")
|
||||||
|
"""
|
||||||
|
from src.domains.agents.plan.agent import (
|
||||||
|
PlanAgentImpl,
|
||||||
|
PlanContext,
|
||||||
|
plan_agent,
|
||||||
|
plan,
|
||||||
|
plan_stream,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"PlanAgentImpl",
|
||||||
|
"PlanContext",
|
||||||
|
"plan_agent",
|
||||||
|
"plan",
|
||||||
|
"plan_stream",
|
||||||
|
]
|
||||||
|
|||||||
@@ -0,0 +1,168 @@
|
|||||||
|
"""
|
||||||
|
Plan Agent implementation using PydanticAI.
|
||||||
|
|
||||||
|
Software architect agent that explores codebases and designs implementation plans.
|
||||||
|
Uses only read-only tools - cannot modify any files.
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from pydantic_ai import Agent
|
||||||
|
from pydantic_ai.models.openai import OpenAIModel
|
||||||
|
|
||||||
|
from src.domains.agents.base import BaseAgent, AgentContext, register_agent
|
||||||
|
from src.domains.agents.plan.prompts import PLAN_SYSTEM_PROMPT
|
||||||
|
from src.ollama.provider import get_ollama_provider
|
||||||
|
from src.shared.config import get_settings
|
||||||
|
from src.shared.logging import logged, get_logger, trace_span
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class PlanContext(AgentContext):
|
||||||
|
"""
|
||||||
|
Context for plan agent tools.
|
||||||
|
|
||||||
|
Passed to all tool functions via RunContext.
|
||||||
|
Uses the same fields as base AgentContext.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class PlanAgentImpl(BaseAgent):
|
||||||
|
"""
|
||||||
|
Software architect agent for implementation planning.
|
||||||
|
|
||||||
|
Explores codebases to understand patterns and conventions,
|
||||||
|
then designs step-by-step implementation plans.
|
||||||
|
|
||||||
|
READ-ONLY: Cannot modify files - uses only exploration tools.
|
||||||
|
"""
|
||||||
|
|
||||||
|
name = "plan"
|
||||||
|
description = "Software architect for designing implementation plans - explores codebase and creates step-by-step strategies"
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
"""Initialize the plan agent."""
|
||||||
|
self._agent: Agent[PlanContext, str] | None = None
|
||||||
|
self._settings = get_settings()
|
||||||
|
|
||||||
|
def _create_agent(self) -> Agent[PlanContext, str]:
|
||||||
|
"""Create the PydanticAI agent with Ollama backend."""
|
||||||
|
# Use sanitized Ollama provider to fix content: null issues
|
||||||
|
model = OpenAIModel(
|
||||||
|
model_name=self._settings.ollama_agent_model,
|
||||||
|
provider=get_ollama_provider(),
|
||||||
|
)
|
||||||
|
|
||||||
|
agent: Agent[PlanContext, str] = Agent(
|
||||||
|
model=model,
|
||||||
|
system_prompt=PLAN_SYSTEM_PROMPT,
|
||||||
|
deps_type=PlanContext,
|
||||||
|
output_type=str,
|
||||||
|
# Mistral Nemo settings:
|
||||||
|
# - temperature 0.3 (Nemo needs slightly higher than 0.0)
|
||||||
|
# - tool_choice "required" forces tool use
|
||||||
|
model_settings={
|
||||||
|
"temperature": 0.3,
|
||||||
|
"extra_body": {"tool_choice": "required"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
# Register read-only tools
|
||||||
|
self._register_tools(agent)
|
||||||
|
|
||||||
|
return agent
|
||||||
|
|
||||||
|
def _register_tools(self, agent: Agent[PlanContext, str]) -> None:
|
||||||
|
"""Register read-only exploration tools with the agent."""
|
||||||
|
from src.domains.agents.plan.tools import register_plan_tools
|
||||||
|
register_plan_tools(agent)
|
||||||
|
|
||||||
|
@logged()
|
||||||
|
async def run(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None,
|
||||||
|
allowed_paths: list[str] | None = None,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Run the plan agent to design an implementation strategy.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt: Description of what to implement
|
||||||
|
working_dir: Working directory for exploration
|
||||||
|
allowed_paths: Restrict tool access to these paths
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Implementation plan with steps and critical files
|
||||||
|
"""
|
||||||
|
ctx = PlanContext(
|
||||||
|
working_dir=working_dir or os.getcwd(),
|
||||||
|
allowed_paths=allowed_paths or self._settings.effective_allowed_paths,
|
||||||
|
timeout_seconds=self._settings.tool_timeout_seconds,
|
||||||
|
)
|
||||||
|
|
||||||
|
async with trace_span("plan_agent_run"):
|
||||||
|
try:
|
||||||
|
# Use run() not run_stream() - Ollama has bugs with streaming + tools
|
||||||
|
result = await self.agent.run(prompt, deps=ctx)
|
||||||
|
return result.output
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"Plan agent error: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def run_stream(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None,
|
||||||
|
allowed_paths: list[str] | None = None,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> AsyncIterator[str]:
|
||||||
|
"""
|
||||||
|
Run the plan agent with streaming output.
|
||||||
|
|
||||||
|
Yields text chunks as they become available.
|
||||||
|
"""
|
||||||
|
ctx = PlanContext(
|
||||||
|
working_dir=working_dir or os.getcwd(),
|
||||||
|
allowed_paths=allowed_paths or self._settings.effective_allowed_paths,
|
||||||
|
timeout_seconds=self._settings.tool_timeout_seconds,
|
||||||
|
)
|
||||||
|
|
||||||
|
async with trace_span("plan_agent_stream"):
|
||||||
|
try:
|
||||||
|
async with self.agent.run_stream(prompt, deps=ctx) as result:
|
||||||
|
async for chunk in result.stream_text():
|
||||||
|
yield chunk
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"Plan agent stream error: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
# Create and register the singleton instance
|
||||||
|
plan_agent = PlanAgentImpl()
|
||||||
|
register_agent(plan_agent)
|
||||||
|
|
||||||
|
|
||||||
|
async def plan(
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> str:
|
||||||
|
"""Run planning query."""
|
||||||
|
return await plan_agent.run(prompt, working_dir=working_dir, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
async def plan_stream(
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> AsyncIterator[str]:
|
||||||
|
"""Run planning query with streaming."""
|
||||||
|
async for chunk in plan_agent.run_stream(prompt, working_dir=working_dir, **kwargs):
|
||||||
|
yield chunk
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
"""
|
||||||
|
System prompts for the Plan agent.
|
||||||
|
|
||||||
|
The Plan agent is a READ-ONLY software architect that explores codebases
|
||||||
|
and designs implementation plans without modifying any files.
|
||||||
|
"""
|
||||||
|
|
||||||
|
PLAN_SYSTEM_PROMPT = """You are a software architect and planning specialist.
|
||||||
|
|
||||||
|
Your role is to explore codebases and design implementation plans.
|
||||||
|
|
||||||
|
CRITICAL: You are READ-ONLY. You CANNOT modify any files.
|
||||||
|
|
||||||
|
AVAILABLE TOOLS:
|
||||||
|
- glob_files: Find files by pattern
|
||||||
|
- read_file: Read file contents
|
||||||
|
- grep_content: Search code with regex
|
||||||
|
- bash_readonly: Run read-only commands (ls, git status, git log, etc.)
|
||||||
|
|
||||||
|
WORKFLOW:
|
||||||
|
1. Understand the requirements
|
||||||
|
2. Explore the codebase to find relevant patterns and conventions
|
||||||
|
3. Design an implementation approach
|
||||||
|
4. Create a step-by-step plan with specific files and changes
|
||||||
|
|
||||||
|
TOOL CALL EXAMPLES (follow exactly):
|
||||||
|
|
||||||
|
To find Python files:
|
||||||
|
Call glob_files with pattern="**/*.py"
|
||||||
|
|
||||||
|
To find a specific file:
|
||||||
|
Call glob_files with pattern="**/config.py"
|
||||||
|
|
||||||
|
To read a file:
|
||||||
|
Call read_file with file_path="/absolute/path/to/file.py"
|
||||||
|
|
||||||
|
To search for code patterns:
|
||||||
|
Call grep_content with pattern="class.*Controller"
|
||||||
|
|
||||||
|
To check git history:
|
||||||
|
Call bash_readonly with command="git log --oneline -10"
|
||||||
|
|
||||||
|
OUTPUT FORMAT:
|
||||||
|
End your response with:
|
||||||
|
|
||||||
|
### Implementation Steps
|
||||||
|
1. [First step with specific file and changes]
|
||||||
|
2. [Second step...]
|
||||||
|
3. [Continue...]
|
||||||
|
|
||||||
|
### Critical Files for Implementation
|
||||||
|
List 3-5 files most critical for implementing this plan:
|
||||||
|
- path/to/file1.py - [Brief reason: e.g., "Core logic to modify"]
|
||||||
|
- path/to/file2.py - [Brief reason: e.g., "Pattern to follow"]
|
||||||
|
|
||||||
|
RULES:
|
||||||
|
- ALWAYS use tools first, then analyze results
|
||||||
|
- Follow existing patterns in the codebase
|
||||||
|
- Consider trade-offs and alternatives
|
||||||
|
- Identify dependencies and sequencing
|
||||||
|
- Never guess - verify with tools
|
||||||
|
- Provide specific file paths and code locations
|
||||||
|
"""
|
||||||
@@ -0,0 +1,170 @@
|
|||||||
|
"""
|
||||||
|
Tool registrations for the Plan agent.
|
||||||
|
|
||||||
|
The Plan agent only has access to READ-ONLY tools.
|
||||||
|
It cannot modify files - only explore and analyze.
|
||||||
|
"""
|
||||||
|
from pydantic_ai import Agent, RunContext
|
||||||
|
|
||||||
|
from src.domains.agents.base import AgentContext
|
||||||
|
from src.domains.tools.file.read import ReadFileTool
|
||||||
|
from src.domains.tools.file.glob import GlobFilesTool
|
||||||
|
from src.domains.tools.search.grep import GrepContentTool
|
||||||
|
from src.domains.tools.shell.bash import BashReadOnlyTool
|
||||||
|
|
||||||
|
|
||||||
|
def register_plan_tools(agent: Agent[AgentContext, str]) -> None:
|
||||||
|
"""
|
||||||
|
Register read-only exploration tools with the Plan agent.
|
||||||
|
|
||||||
|
The Plan agent is restricted to read-only tools:
|
||||||
|
- read_file: Read file contents
|
||||||
|
- glob_files: Find files by pattern
|
||||||
|
- grep_content: Search file contents
|
||||||
|
- bash_readonly: Read-only shell commands
|
||||||
|
|
||||||
|
Write tools (edit_file, write_file, bash) are NOT available.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def read_file(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
file_path: str,
|
||||||
|
offset: int = 0,
|
||||||
|
limit: int = 2000
|
||||||
|
) -> str:
|
||||||
|
"""Read contents of a file with line numbers.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_path: Absolute path to the file to read
|
||||||
|
offset: Line number to start from (0-based, default: 0)
|
||||||
|
limit: Maximum number of lines to read (default: 2000)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
File contents with line numbers, or error message.
|
||||||
|
|
||||||
|
IMPORTANT: Always use absolute paths. Use this to understand existing code.
|
||||||
|
"""
|
||||||
|
tool = ReadFileTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
result = await tool.execute(
|
||||||
|
file_path=file_path,
|
||||||
|
offset=offset,
|
||||||
|
limit=limit
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def glob_files(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
pattern: str,
|
||||||
|
path: str | None = None,
|
||||||
|
limit: int = 100
|
||||||
|
) -> str:
|
||||||
|
"""Find files matching a glob pattern.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pattern: Glob pattern (e.g., "**/*.py", "src/**/*.ts", "*.md")
|
||||||
|
path: Directory to search in (default: working directory)
|
||||||
|
limit: Maximum number of files to return (default: 100)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of absolute file paths, sorted by modification time (newest first).
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
- "**/*.py" finds all Python files
|
||||||
|
- "src/**/*.ts" finds TypeScript files in src/
|
||||||
|
- "**/test_*.py" finds all test files
|
||||||
|
|
||||||
|
IMPORTANT: Use this to discover files before reading them.
|
||||||
|
"""
|
||||||
|
tool = GlobFilesTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
search_path = path or ctx.deps.working_dir
|
||||||
|
result = await tool.execute(
|
||||||
|
pattern=pattern,
|
||||||
|
path=search_path,
|
||||||
|
limit=limit
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def grep_content(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
pattern: str,
|
||||||
|
path: str | None = None,
|
||||||
|
file_glob: str | None = None,
|
||||||
|
context_lines: int = 0,
|
||||||
|
case_sensitive: bool = True
|
||||||
|
) -> str:
|
||||||
|
"""Search file contents using regex pattern.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pattern: Regex pattern to search for (Python re syntax)
|
||||||
|
path: Directory or file to search (default: working directory)
|
||||||
|
file_glob: Filter files by glob (e.g., "*.py", "*.ts")
|
||||||
|
context_lines: Lines of context before/after matches (default: 0)
|
||||||
|
case_sensitive: Case-sensitive search (default: True)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Matching lines with file paths and line numbers.
|
||||||
|
Format: "filepath:line_num: content"
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
- pattern="def.*__init__" finds init methods
|
||||||
|
- pattern="class\\s+\\w+" finds class definitions
|
||||||
|
- pattern="TODO|FIXME" finds todo comments
|
||||||
|
|
||||||
|
IMPORTANT: Use this to find code patterns and implementations.
|
||||||
|
"""
|
||||||
|
tool = GrepContentTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
search_path = path or ctx.deps.working_dir
|
||||||
|
result = await tool.execute(
|
||||||
|
pattern=pattern,
|
||||||
|
path=search_path,
|
||||||
|
file_glob=file_glob,
|
||||||
|
context_lines=context_lines,
|
||||||
|
case_sensitive=case_sensitive
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def bash_readonly(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
command: str,
|
||||||
|
cwd: str | None = None,
|
||||||
|
timeout: int = 30
|
||||||
|
) -> str:
|
||||||
|
"""Execute a read-only bash command.
|
||||||
|
|
||||||
|
ALLOWED commands:
|
||||||
|
- File inspection: ls, find, cat, head, tail, wc, file, stat, tree, du
|
||||||
|
- Git (read-only): git status, git log, git diff, git show, git branch
|
||||||
|
- Text processing: grep, awk, sed (read-only), sort, uniq
|
||||||
|
- System info: pwd, whoami, hostname, which
|
||||||
|
|
||||||
|
FORBIDDEN:
|
||||||
|
- File modification (rm, mv, cp, mkdir, touch)
|
||||||
|
- Redirects (>, >>)
|
||||||
|
- Command chaining (&&, ||, ;)
|
||||||
|
- Network (curl, wget)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
command: The bash command to execute
|
||||||
|
cwd: Working directory (default: agent working directory)
|
||||||
|
timeout: Timeout in seconds (default: 30)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Command output or error message.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
- "ls -la" lists files with details
|
||||||
|
- "git status" shows git status
|
||||||
|
- "git log --oneline -10" shows recent commits
|
||||||
|
"""
|
||||||
|
tool = BashReadOnlyTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
working_dir = cwd or ctx.deps.working_dir
|
||||||
|
result = await tool.execute(
|
||||||
|
command=command,
|
||||||
|
cwd=working_dir,
|
||||||
|
timeout=min(timeout, ctx.deps.timeout_seconds)
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
@@ -9,6 +9,8 @@ from src.domains.agents.base import get_agent, list_agents
|
|||||||
|
|
||||||
# Import agents to ensure they're registered
|
# Import agents to ensure they're registered
|
||||||
import src.domains.agents.explore # noqa: F401
|
import src.domains.agents.explore # noqa: F401
|
||||||
|
import src.domains.agents.plan # noqa: F401
|
||||||
|
import src.domains.agents.task # noqa: F401
|
||||||
from src.domains.agents.schemas import (
|
from src.domains.agents.schemas import (
|
||||||
AgentRunRequest,
|
AgentRunRequest,
|
||||||
AgentRunResponse,
|
AgentRunResponse,
|
||||||
|
|||||||
@@ -0,0 +1,33 @@
|
|||||||
|
"""
|
||||||
|
Task Agent - Full orchestrator for autonomous task execution.
|
||||||
|
|
||||||
|
The Task agent can:
|
||||||
|
- Execute multi-step tasks autonomously
|
||||||
|
- Use all tools (read + write + bash)
|
||||||
|
- Spawn sub-agents (Explore, Plan) for focused work
|
||||||
|
- Return consolidated task summaries
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
from src.domains.agents.task import task_agent, task
|
||||||
|
|
||||||
|
# Direct agent access
|
||||||
|
result = await task_agent.run("Create a new user model with tests")
|
||||||
|
|
||||||
|
# Convenience function
|
||||||
|
result = await task("Create a new user model with tests")
|
||||||
|
"""
|
||||||
|
from src.domains.agents.task.agent import (
|
||||||
|
TaskAgentImpl,
|
||||||
|
TaskContext,
|
||||||
|
task_agent,
|
||||||
|
task,
|
||||||
|
task_stream,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"TaskAgentImpl",
|
||||||
|
"TaskContext",
|
||||||
|
"task_agent",
|
||||||
|
"task",
|
||||||
|
"task_stream",
|
||||||
|
]
|
||||||
|
|||||||
@@ -0,0 +1,174 @@
|
|||||||
|
"""
|
||||||
|
Task Agent implementation using PydanticAI.
|
||||||
|
|
||||||
|
Full orchestrator agent that can:
|
||||||
|
- Execute multi-step tasks autonomously
|
||||||
|
- Use all tools (read + write)
|
||||||
|
- Spawn sub-agents (Explore, Plan) for focused work
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from pydantic_ai import Agent
|
||||||
|
from pydantic_ai.models.openai import OpenAIModel
|
||||||
|
|
||||||
|
from src.domains.agents.base import BaseAgent, AgentContext, register_agent
|
||||||
|
from src.domains.agents.task.prompts import TASK_SYSTEM_PROMPT
|
||||||
|
from src.ollama.provider import get_ollama_provider
|
||||||
|
from src.shared.config import get_settings
|
||||||
|
from src.shared.logging import logged, get_logger, trace_span
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TaskContext(AgentContext):
|
||||||
|
"""
|
||||||
|
Context for task agent tools.
|
||||||
|
|
||||||
|
Passed to all tool functions via RunContext.
|
||||||
|
Uses the same fields as base AgentContext.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class TaskAgentImpl(BaseAgent):
|
||||||
|
"""
|
||||||
|
Full orchestrator agent for autonomous task execution.
|
||||||
|
|
||||||
|
Has access to ALL tools:
|
||||||
|
- Read-only: read_file, glob_files, grep_content, bash_readonly
|
||||||
|
- Write: edit_file, write_file, bash
|
||||||
|
- External: web_search
|
||||||
|
- Orchestration: spawn_agent (launch sub-agents)
|
||||||
|
|
||||||
|
Can spawn Explore and Plan agents to offload focused tasks,
|
||||||
|
keeping context efficient across complex multi-step work.
|
||||||
|
"""
|
||||||
|
|
||||||
|
name = "task"
|
||||||
|
description = "Autonomous multi-step task execution with sub-agent orchestration"
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
"""Initialize the task agent."""
|
||||||
|
self._agent: Agent[TaskContext, str] | None = None
|
||||||
|
self._settings = get_settings()
|
||||||
|
|
||||||
|
def _create_agent(self) -> Agent[TaskContext, str]:
|
||||||
|
"""Create the PydanticAI agent with Ollama backend."""
|
||||||
|
# Use sanitized Ollama provider to fix content: null issues
|
||||||
|
model = OpenAIModel(
|
||||||
|
model_name=self._settings.ollama_agent_model,
|
||||||
|
provider=get_ollama_provider(),
|
||||||
|
)
|
||||||
|
|
||||||
|
agent: Agent[TaskContext, str] = Agent(
|
||||||
|
model=model,
|
||||||
|
system_prompt=TASK_SYSTEM_PROMPT,
|
||||||
|
deps_type=TaskContext,
|
||||||
|
output_type=str,
|
||||||
|
# Mistral Nemo settings:
|
||||||
|
# - temperature 0.3 (Nemo needs slightly higher than 0.0)
|
||||||
|
# - tool_choice "required" forces tool use
|
||||||
|
model_settings={
|
||||||
|
"temperature": 0.3,
|
||||||
|
"extra_body": {"tool_choice": "required"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
# Register all tools including orchestration
|
||||||
|
self._register_tools(agent)
|
||||||
|
|
||||||
|
return agent
|
||||||
|
|
||||||
|
def _register_tools(self, agent: Agent[TaskContext, str]) -> None:
|
||||||
|
"""Register all tools with the agent."""
|
||||||
|
from src.domains.agents.task.tools import register_task_tools
|
||||||
|
register_task_tools(agent)
|
||||||
|
|
||||||
|
@logged()
|
||||||
|
async def run(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None,
|
||||||
|
allowed_paths: list[str] | None = None,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Run the task agent to execute a multi-step task.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt: Description of the task to execute
|
||||||
|
working_dir: Working directory for the agent
|
||||||
|
allowed_paths: Restrict tool access to these paths
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Consolidated task summary with results
|
||||||
|
"""
|
||||||
|
ctx = TaskContext(
|
||||||
|
working_dir=working_dir or os.getcwd(),
|
||||||
|
allowed_paths=allowed_paths or self._settings.effective_allowed_paths,
|
||||||
|
timeout_seconds=self._settings.tool_timeout_seconds,
|
||||||
|
)
|
||||||
|
|
||||||
|
async with trace_span("task_agent_run"):
|
||||||
|
try:
|
||||||
|
# Use run() not run_stream() - Ollama has bugs with streaming + tools
|
||||||
|
result = await self.agent.run(prompt, deps=ctx)
|
||||||
|
return result.output
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"Task agent error: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def run_stream(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None,
|
||||||
|
allowed_paths: list[str] | None = None,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> AsyncIterator[str]:
|
||||||
|
"""
|
||||||
|
Run the task agent with streaming output.
|
||||||
|
|
||||||
|
Yields text chunks as they become available.
|
||||||
|
"""
|
||||||
|
ctx = TaskContext(
|
||||||
|
working_dir=working_dir or os.getcwd(),
|
||||||
|
allowed_paths=allowed_paths or self._settings.effective_allowed_paths,
|
||||||
|
timeout_seconds=self._settings.tool_timeout_seconds,
|
||||||
|
)
|
||||||
|
|
||||||
|
async with trace_span("task_agent_stream"):
|
||||||
|
try:
|
||||||
|
async with self.agent.run_stream(prompt, deps=ctx) as result:
|
||||||
|
async for chunk in result.stream_text():
|
||||||
|
yield chunk
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"Task agent stream error: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
# Create and register the singleton instance
|
||||||
|
task_agent = TaskAgentImpl()
|
||||||
|
register_agent(task_agent)
|
||||||
|
|
||||||
|
|
||||||
|
async def task(
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> str:
|
||||||
|
"""Run task execution."""
|
||||||
|
return await task_agent.run(prompt, working_dir=working_dir, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
async def task_stream(
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> AsyncIterator[str]:
|
||||||
|
"""Run task execution with streaming."""
|
||||||
|
async for chunk in task_agent.run_stream(prompt, working_dir=working_dir, **kwargs):
|
||||||
|
yield chunk
|
||||||
@@ -0,0 +1,83 @@
|
|||||||
|
"""
|
||||||
|
System prompts for the Task agent.
|
||||||
|
|
||||||
|
The Task agent is a full orchestrator that can:
|
||||||
|
- Execute multi-step tasks autonomously
|
||||||
|
- Use all tools (read + write)
|
||||||
|
- Spawn sub-agents (Explore, Plan) for focused work
|
||||||
|
"""
|
||||||
|
|
||||||
|
TASK_SYSTEM_PROMPT = """You are an autonomous task execution agent.
|
||||||
|
|
||||||
|
You have access to ALL tools including file editing, writing, and bash execution.
|
||||||
|
You can also spawn sub-agents to help with complex tasks.
|
||||||
|
|
||||||
|
AVAILABLE TOOLS:
|
||||||
|
|
||||||
|
File Operations:
|
||||||
|
- read_file: Read file contents with line numbers
|
||||||
|
- glob_files: Find files by pattern
|
||||||
|
- grep_content: Search file contents with regex
|
||||||
|
- edit_file: Make targeted edits via find-and-replace
|
||||||
|
- write_file: Create or overwrite files
|
||||||
|
|
||||||
|
Shell:
|
||||||
|
- bash_readonly: Read-only commands (ls, git status, git log, etc.)
|
||||||
|
- bash: Full bash execution (git commit, pytest, mkdir, etc.)
|
||||||
|
|
||||||
|
External:
|
||||||
|
- web_search: Search the web for current information
|
||||||
|
|
||||||
|
Orchestration:
|
||||||
|
- spawn_agent: Launch sub-agents for focused tasks
|
||||||
|
|
||||||
|
WORKFLOW:
|
||||||
|
1. Understand the task requirements
|
||||||
|
2. Break down into sub-tasks if complex
|
||||||
|
3. Use spawn_agent for research (explore) or planning (plan)
|
||||||
|
4. Execute implementation steps using write tools
|
||||||
|
5. Validate changes (run tests if applicable)
|
||||||
|
6. Return consolidated summary
|
||||||
|
|
||||||
|
TOOL CALL EXAMPLES:
|
||||||
|
|
||||||
|
To spawn an Explore agent for research:
|
||||||
|
Call spawn_agent with agent_type="explore" and prompt="find all config files"
|
||||||
|
|
||||||
|
To spawn a Plan agent for design:
|
||||||
|
Call spawn_agent with agent_type="plan" and prompt="design user auth feature"
|
||||||
|
|
||||||
|
To edit a file:
|
||||||
|
Call edit_file with file_path="/path/to/file.py" and old_string="old" and new_string="new"
|
||||||
|
|
||||||
|
To run tests:
|
||||||
|
Call bash with command="pytest tests/ -v"
|
||||||
|
|
||||||
|
SPAWN_AGENT USAGE:
|
||||||
|
- Use spawn_agent to offload focused tasks to specialized agents
|
||||||
|
- Explore agent: Fast codebase searches and analysis
|
||||||
|
- Plan agent: Design implementation strategies
|
||||||
|
- Keep each agent's context focused and efficient
|
||||||
|
|
||||||
|
GIT DISCIPLINE:
|
||||||
|
- Create feature branches for changes
|
||||||
|
- Use conventional commit format (feat:, fix:, docs:, etc.)
|
||||||
|
- Never commit directly to main
|
||||||
|
- Run tests before committing
|
||||||
|
|
||||||
|
RULES:
|
||||||
|
- ALWAYS use tools first, then analyze results
|
||||||
|
- Never guess file contents - read them first
|
||||||
|
- Prefer edit_file over write_file for existing files
|
||||||
|
- Use spawn_agent to keep context focused
|
||||||
|
- Validate changes by running tests when applicable
|
||||||
|
|
||||||
|
OUTPUT FORMAT:
|
||||||
|
End your response with a summary:
|
||||||
|
|
||||||
|
### Task Summary
|
||||||
|
- **Accomplished:** What was done
|
||||||
|
- **Files modified:** List of changed files
|
||||||
|
- **Commands run:** Key commands executed
|
||||||
|
- **Issues:** Any problems encountered
|
||||||
|
"""
|
||||||
@@ -0,0 +1,350 @@
|
|||||||
|
"""
|
||||||
|
Tool registrations for the Task agent.
|
||||||
|
|
||||||
|
The Task agent has access to ALL tools:
|
||||||
|
- Read-only tools (same as Explore/Plan)
|
||||||
|
- Write tools (edit, write, bash full)
|
||||||
|
- External tools (web search)
|
||||||
|
- Orchestration (spawn sub-agents)
|
||||||
|
"""
|
||||||
|
from pydantic_ai import Agent, RunContext
|
||||||
|
|
||||||
|
from src.domains.agents.base import AgentContext
|
||||||
|
from src.domains.tools.file.read import ReadFileTool
|
||||||
|
from src.domains.tools.file.glob import GlobFilesTool
|
||||||
|
from src.domains.tools.file.edit import EditFileTool
|
||||||
|
from src.domains.tools.file.write import WriteFileTool
|
||||||
|
from src.domains.tools.search.grep import GrepContentTool
|
||||||
|
from src.domains.tools.search.web import WebSearchTool
|
||||||
|
from src.domains.tools.shell.bash import BashReadOnlyTool
|
||||||
|
from src.domains.tools.shell.bash_full import BashTool
|
||||||
|
|
||||||
|
|
||||||
|
def register_task_tools(agent: Agent[AgentContext, str]) -> None:
|
||||||
|
"""
|
||||||
|
Register all tools with the Task agent.
|
||||||
|
|
||||||
|
Includes:
|
||||||
|
- Read-only tools: read_file, glob_files, grep_content, bash_readonly
|
||||||
|
- Write tools: edit_file, write_file, bash
|
||||||
|
- External: web_search
|
||||||
|
- Orchestration: spawn_agent
|
||||||
|
"""
|
||||||
|
|
||||||
|
# === Read-only tools ===
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def read_file(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
file_path: str,
|
||||||
|
offset: int = 0,
|
||||||
|
limit: int = 2000
|
||||||
|
) -> str:
|
||||||
|
"""Read contents of a file with line numbers.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_path: Absolute path to the file to read
|
||||||
|
offset: Line number to start from (0-based, default: 0)
|
||||||
|
limit: Maximum number of lines to read (default: 2000)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
File contents with line numbers, or error message.
|
||||||
|
|
||||||
|
IMPORTANT: Always use absolute paths. Read files before editing them.
|
||||||
|
"""
|
||||||
|
tool = ReadFileTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
result = await tool.execute(
|
||||||
|
file_path=file_path,
|
||||||
|
offset=offset,
|
||||||
|
limit=limit
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def glob_files(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
pattern: str,
|
||||||
|
path: str | None = None,
|
||||||
|
limit: int = 100
|
||||||
|
) -> str:
|
||||||
|
"""Find files matching a glob pattern.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pattern: Glob pattern (e.g., "**/*.py", "src/**/*.ts", "*.md")
|
||||||
|
path: Directory to search in (default: working directory)
|
||||||
|
limit: Maximum number of files to return (default: 100)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of absolute file paths, sorted by modification time (newest first).
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
- "**/*.py" finds all Python files
|
||||||
|
- "src/**/*.ts" finds TypeScript files in src/
|
||||||
|
- "**/test_*.py" finds all test files
|
||||||
|
"""
|
||||||
|
tool = GlobFilesTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
search_path = path or ctx.deps.working_dir
|
||||||
|
result = await tool.execute(
|
||||||
|
pattern=pattern,
|
||||||
|
path=search_path,
|
||||||
|
limit=limit
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def grep_content(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
pattern: str,
|
||||||
|
path: str | None = None,
|
||||||
|
file_glob: str | None = None,
|
||||||
|
context_lines: int = 0,
|
||||||
|
case_sensitive: bool = True
|
||||||
|
) -> str:
|
||||||
|
"""Search file contents using regex pattern.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pattern: Regex pattern to search for (Python re syntax)
|
||||||
|
path: Directory or file to search (default: working directory)
|
||||||
|
file_glob: Filter files by glob (e.g., "*.py", "*.ts")
|
||||||
|
context_lines: Lines of context before/after matches (default: 0)
|
||||||
|
case_sensitive: Case-sensitive search (default: True)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Matching lines with file paths and line numbers.
|
||||||
|
Format: "filepath:line_num: content"
|
||||||
|
"""
|
||||||
|
tool = GrepContentTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
search_path = path or ctx.deps.working_dir
|
||||||
|
result = await tool.execute(
|
||||||
|
pattern=pattern,
|
||||||
|
path=search_path,
|
||||||
|
file_glob=file_glob,
|
||||||
|
context_lines=context_lines,
|
||||||
|
case_sensitive=case_sensitive
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def bash_readonly(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
command: str,
|
||||||
|
cwd: str | None = None,
|
||||||
|
timeout: int = 30
|
||||||
|
) -> str:
|
||||||
|
"""Execute a read-only bash command.
|
||||||
|
|
||||||
|
ALLOWED commands:
|
||||||
|
- File inspection: ls, find, cat, head, tail, wc, file, stat, tree, du
|
||||||
|
- Git (read-only): git status, git log, git diff, git show, git branch
|
||||||
|
- Text processing: grep, awk, sed (read-only), sort, uniq
|
||||||
|
- System info: pwd, whoami, hostname, which
|
||||||
|
|
||||||
|
FORBIDDEN:
|
||||||
|
- File modification (rm, mv, cp, mkdir, touch)
|
||||||
|
- Redirects (>, >>)
|
||||||
|
- Command chaining (&&, ||, ;)
|
||||||
|
- Network (curl, wget)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
command: The bash command to execute
|
||||||
|
cwd: Working directory (default: agent working directory)
|
||||||
|
timeout: Timeout in seconds (default: 30)
|
||||||
|
"""
|
||||||
|
tool = BashReadOnlyTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
working_dir = cwd or ctx.deps.working_dir
|
||||||
|
result = await tool.execute(
|
||||||
|
command=command,
|
||||||
|
cwd=working_dir,
|
||||||
|
timeout=min(timeout, ctx.deps.timeout_seconds)
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
# === Write tools ===
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def edit_file(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
file_path: str,
|
||||||
|
old_string: str,
|
||||||
|
new_string: str,
|
||||||
|
replace_all: bool = False
|
||||||
|
) -> str:
|
||||||
|
"""Make targeted edits to a file using find-and-replace.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_path: Absolute path to the file to edit
|
||||||
|
old_string: The exact text to find and replace (must exist in file)
|
||||||
|
new_string: The replacement text
|
||||||
|
replace_all: If True, replace all occurrences. If False (default),
|
||||||
|
old_string must be unique (appear exactly once).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Success message with diff preview, or error.
|
||||||
|
|
||||||
|
IMPORTANT:
|
||||||
|
- old_string must exactly match file content (including whitespace)
|
||||||
|
- By default, old_string must appear exactly once (for safety)
|
||||||
|
- Always read the file first to verify exact content before editing
|
||||||
|
"""
|
||||||
|
tool = EditFileTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
result = await tool.execute(
|
||||||
|
file_path=file_path,
|
||||||
|
old_string=old_string,
|
||||||
|
new_string=new_string,
|
||||||
|
replace_all=replace_all
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def write_file(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
file_path: str,
|
||||||
|
content: str
|
||||||
|
) -> str:
|
||||||
|
"""Create a new file or overwrite an existing file.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_path: Absolute path to the file to create/write
|
||||||
|
content: The content to write to the file
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Success message with file path and size.
|
||||||
|
|
||||||
|
IMPORTANT:
|
||||||
|
- Parent directory must exist (use bash mkdir first if needed)
|
||||||
|
- For editing existing files, prefer edit_file instead
|
||||||
|
- Will overwrite existing files without confirmation
|
||||||
|
"""
|
||||||
|
tool = WriteFileTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
result = await tool.execute(
|
||||||
|
file_path=file_path,
|
||||||
|
content=content
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def bash(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
command: str,
|
||||||
|
cwd: str | None = None,
|
||||||
|
timeout: int = 60
|
||||||
|
) -> str:
|
||||||
|
"""Execute a bash command with write capabilities.
|
||||||
|
|
||||||
|
ALLOWED:
|
||||||
|
- File operations: ls, find, mkdir, touch, cp, mv, rm (single files)
|
||||||
|
- Git (full): git add, git commit, git checkout, git merge, git pull
|
||||||
|
- Python: python, pip install, pytest, mypy, ruff
|
||||||
|
- Text processing: grep, awk, sed, sort
|
||||||
|
- Command chaining: && and || are allowed
|
||||||
|
|
||||||
|
FORBIDDEN:
|
||||||
|
- sudo, su (privilege escalation)
|
||||||
|
- Network: curl, wget, ssh, scp, rsync
|
||||||
|
- Dangerous: rm -rf, chmod 777, dd, mkfs
|
||||||
|
|
||||||
|
Args:
|
||||||
|
command: The bash command to execute
|
||||||
|
cwd: Working directory (default: agent working directory)
|
||||||
|
timeout: Timeout in seconds (default: 60)
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
- "mkdir -p src/utils" creates directory
|
||||||
|
- "git add . && git commit -m 'fix: bug'" commits changes
|
||||||
|
- "pytest tests/ -v" runs tests
|
||||||
|
"""
|
||||||
|
tool = BashTool(allowed_paths=ctx.deps.allowed_paths)
|
||||||
|
working_dir = cwd or ctx.deps.working_dir
|
||||||
|
result = await tool.execute(
|
||||||
|
command=command,
|
||||||
|
cwd=working_dir,
|
||||||
|
timeout=min(timeout, ctx.deps.timeout_seconds)
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
# === External tools ===
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def web_search(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
query: str,
|
||||||
|
num_results: int = 5,
|
||||||
|
categories: str | None = None
|
||||||
|
) -> str:
|
||||||
|
"""Search the web for current information.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query: Search query (e.g., "Python 3.12 new features")
|
||||||
|
num_results: Number of results to return (1-10, default: 5)
|
||||||
|
categories: Optional category filter ("general", "it", "news", "science")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Search results with titles, URLs, and snippets.
|
||||||
|
|
||||||
|
Use this for:
|
||||||
|
- Current events or recent information
|
||||||
|
- Documentation updates
|
||||||
|
- Technical references with URLs
|
||||||
|
"""
|
||||||
|
tool = WebSearchTool()
|
||||||
|
result = await tool.execute(
|
||||||
|
query=query,
|
||||||
|
num_results=num_results,
|
||||||
|
categories=categories
|
||||||
|
)
|
||||||
|
return result.to_string()
|
||||||
|
|
||||||
|
# === Orchestration tools ===
|
||||||
|
|
||||||
|
@agent.tool
|
||||||
|
async def spawn_agent(
|
||||||
|
ctx: RunContext[AgentContext],
|
||||||
|
agent_type: str,
|
||||||
|
prompt: str,
|
||||||
|
working_dir: str | None = None
|
||||||
|
) -> str:
|
||||||
|
"""Spawn a sub-agent to handle a focused task.
|
||||||
|
|
||||||
|
Use this to offload work to specialized agents:
|
||||||
|
- "explore": Fast codebase searches and analysis (read-only)
|
||||||
|
- "plan": Design implementation strategies (read-only)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
agent_type: Type of agent to spawn ("explore" or "plan")
|
||||||
|
prompt: Task description for the sub-agent
|
||||||
|
working_dir: Working directory for the sub-agent (default: current)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Sub-agent's consolidated response.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
- spawn_agent(agent_type="explore", prompt="find all test files")
|
||||||
|
- spawn_agent(agent_type="plan", prompt="design user auth feature")
|
||||||
|
|
||||||
|
IMPORTANT:
|
||||||
|
- Use sub-agents to keep context focused and efficient
|
||||||
|
- Explore agent for research, Plan agent for design
|
||||||
|
- Cannot spawn nested Task agents (recursion risk)
|
||||||
|
"""
|
||||||
|
from src.domains.agents.base import get_agent
|
||||||
|
|
||||||
|
# Validate agent type
|
||||||
|
allowed_types = ["explore", "plan"]
|
||||||
|
if agent_type not in allowed_types:
|
||||||
|
if agent_type == "task":
|
||||||
|
return "Error: Cannot spawn nested Task agents (recursion risk)"
|
||||||
|
return f"Error: Unknown agent type '{agent_type}'. Allowed: {allowed_types}"
|
||||||
|
|
||||||
|
sub_agent = get_agent(agent_type)
|
||||||
|
if not sub_agent:
|
||||||
|
return f"Error: Agent '{agent_type}' not found in registry"
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = await sub_agent.run(
|
||||||
|
prompt=prompt,
|
||||||
|
working_dir=working_dir or ctx.deps.working_dir,
|
||||||
|
allowed_paths=ctx.deps.allowed_paths,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
except Exception as e:
|
||||||
|
return f"Sub-agent error: {e}"
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
"""
|
||||||
|
Conversations domain - Multi-turn conversation management.
|
||||||
|
|
||||||
|
Provides:
|
||||||
|
- Conversation persistence with message history
|
||||||
|
- Context summarization when approaching token limits
|
||||||
|
- Agent integration with conversation context injection
|
||||||
|
"""
|
||||||
|
from src.domains.conversations.models import Conversation, Message
|
||||||
|
from src.domains.conversations.service import ConversationService
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Conversation",
|
||||||
|
"Message",
|
||||||
|
"ConversationService",
|
||||||
|
]
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
"""
|
||||||
|
Database models for conversations.
|
||||||
|
|
||||||
|
Following core-api patterns: SQLAlchemy 2.0 with async support.
|
||||||
|
"""
|
||||||
|
from datetime import datetime
|
||||||
|
from uuid import UUID, uuid4
|
||||||
|
|
||||||
|
from sqlalchemy import ForeignKey, String, Text
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||||
|
|
||||||
|
from src.db.models import Base
|
||||||
|
|
||||||
|
|
||||||
|
class Conversation(Base):
|
||||||
|
"""
|
||||||
|
A conversation session with an agent.
|
||||||
|
|
||||||
|
Tracks message history, token usage, and metadata.
|
||||||
|
"""
|
||||||
|
__tablename__ = "conversations"
|
||||||
|
|
||||||
|
id: Mapped[UUID] = mapped_column(primary_key=True, default=uuid4)
|
||||||
|
user_id: Mapped[str] = mapped_column(String(255), index=True)
|
||||||
|
agent_type: Mapped[str] = mapped_column(String(50), default="explore", insert_default="explore")
|
||||||
|
title: Mapped[str | None] = mapped_column(String(255), nullable=True, default=None)
|
||||||
|
working_dir: Mapped[str] = mapped_column(String(1024), default=".", insert_default=".")
|
||||||
|
total_tokens: Mapped[int] = mapped_column(default=0, insert_default=0)
|
||||||
|
created_at: Mapped[datetime] = mapped_column(default=datetime.utcnow)
|
||||||
|
updated_at: Mapped[datetime | None] = mapped_column(
|
||||||
|
default=datetime.utcnow,
|
||||||
|
onupdate=datetime.utcnow,
|
||||||
|
nullable=True
|
||||||
|
)
|
||||||
|
|
||||||
|
# Relationships
|
||||||
|
messages: Mapped[list["Message"]] = relationship(
|
||||||
|
back_populates="conversation",
|
||||||
|
cascade="all, delete-orphan",
|
||||||
|
order_by="Message.created_at",
|
||||||
|
)
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
return f"<Conversation {self.id} agent={self.agent_type}>"
|
||||||
|
|
||||||
|
|
||||||
|
class Message(Base):
|
||||||
|
"""
|
||||||
|
A single message in a conversation.
|
||||||
|
|
||||||
|
Tracks role, content, token count, and summarization state.
|
||||||
|
"""
|
||||||
|
__tablename__ = "messages"
|
||||||
|
|
||||||
|
id: Mapped[UUID] = mapped_column(primary_key=True, default=uuid4)
|
||||||
|
conversation_id: Mapped[UUID] = mapped_column(
|
||||||
|
ForeignKey("conversations.id", ondelete="CASCADE"),
|
||||||
|
index=True
|
||||||
|
)
|
||||||
|
role: Mapped[str] = mapped_column(String(20)) # user, assistant, system, summary
|
||||||
|
content: Mapped[str] = mapped_column(Text)
|
||||||
|
token_count: Mapped[int] = mapped_column(default=0, insert_default=0)
|
||||||
|
is_summary: Mapped[bool] = mapped_column(default=False, insert_default=False)
|
||||||
|
summarizes_up_to: Mapped[UUID | None] = mapped_column(nullable=True, default=None)
|
||||||
|
created_at: Mapped[datetime] = mapped_column(default=datetime.utcnow)
|
||||||
|
|
||||||
|
# Relationships
|
||||||
|
conversation: Mapped["Conversation"] = relationship(back_populates="messages")
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
preview = self.content[:30] + "..." if len(self.content) > 30 else self.content
|
||||||
|
return f"<Message {self.role}: {preview}>"
|
||||||
@@ -0,0 +1,189 @@
|
|||||||
|
"""
|
||||||
|
REST API routes for conversations.
|
||||||
|
"""
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.db import get_session
|
||||||
|
from src.domains.conversations.schemas import (
|
||||||
|
AddMessageRequest,
|
||||||
|
AddMessageResponse,
|
||||||
|
ConversationDetailResponse,
|
||||||
|
ConversationListResponse,
|
||||||
|
ConversationResponse,
|
||||||
|
CreateConversationRequest,
|
||||||
|
MessageResponse,
|
||||||
|
)
|
||||||
|
from src.domains.conversations.service import ConversationService
|
||||||
|
from src.shared.auth import require_auth
|
||||||
|
from src.shared.logging import logged, get_logger
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/conversations", tags=["Conversations"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/", response_model=ConversationResponse, status_code=201)
|
||||||
|
@logged()
|
||||||
|
async def create_conversation(
|
||||||
|
request: CreateConversationRequest,
|
||||||
|
session: AsyncSession = Depends(get_session),
|
||||||
|
user=Depends(require_auth),
|
||||||
|
) -> ConversationResponse:
|
||||||
|
"""
|
||||||
|
Create a new conversation.
|
||||||
|
|
||||||
|
Starts an empty conversation with the specified agent type.
|
||||||
|
"""
|
||||||
|
service = ConversationService(session)
|
||||||
|
conversation = await service.create(
|
||||||
|
user_id=user.id,
|
||||||
|
agent_type=request.agent_type,
|
||||||
|
working_dir=request.working_dir,
|
||||||
|
title=request.title,
|
||||||
|
)
|
||||||
|
return ConversationResponse.model_validate(conversation)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/", response_model=ConversationListResponse)
|
||||||
|
@logged()
|
||||||
|
async def list_conversations(
|
||||||
|
limit: int = 50,
|
||||||
|
offset: int = 0,
|
||||||
|
session: AsyncSession = Depends(get_session),
|
||||||
|
user=Depends(require_auth),
|
||||||
|
) -> ConversationListResponse:
|
||||||
|
"""
|
||||||
|
List user's conversations.
|
||||||
|
|
||||||
|
Returns conversations sorted by most recently updated.
|
||||||
|
"""
|
||||||
|
service = ConversationService(session)
|
||||||
|
conversations, total = await service.list_by_user(
|
||||||
|
user_id=user.id,
|
||||||
|
limit=limit,
|
||||||
|
offset=offset,
|
||||||
|
)
|
||||||
|
return ConversationListResponse(
|
||||||
|
conversations=[ConversationResponse.model_validate(c) for c in conversations],
|
||||||
|
total=total,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{conversation_id}", response_model=ConversationDetailResponse)
|
||||||
|
@logged()
|
||||||
|
async def get_conversation(
|
||||||
|
conversation_id: UUID,
|
||||||
|
session: AsyncSession = Depends(get_session),
|
||||||
|
user=Depends(require_auth),
|
||||||
|
) -> ConversationDetailResponse:
|
||||||
|
"""
|
||||||
|
Get conversation with all messages.
|
||||||
|
|
||||||
|
Returns conversation metadata and full message history.
|
||||||
|
"""
|
||||||
|
service = ConversationService(session)
|
||||||
|
conversation = await service.get_with_messages(conversation_id)
|
||||||
|
|
||||||
|
if not conversation:
|
||||||
|
raise HTTPException(status_code=404, detail="Conversation not found")
|
||||||
|
|
||||||
|
if conversation.user_id != user.id:
|
||||||
|
raise HTTPException(status_code=403, detail="Not authorized")
|
||||||
|
|
||||||
|
return ConversationDetailResponse.model_validate(conversation)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/{conversation_id}", status_code=204)
|
||||||
|
@logged()
|
||||||
|
async def delete_conversation(
|
||||||
|
conversation_id: UUID,
|
||||||
|
session: AsyncSession = Depends(get_session),
|
||||||
|
user=Depends(require_auth),
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Delete a conversation and all its messages.
|
||||||
|
"""
|
||||||
|
service = ConversationService(session)
|
||||||
|
conversation = await service.get(conversation_id)
|
||||||
|
|
||||||
|
if not conversation:
|
||||||
|
raise HTTPException(status_code=404, detail="Conversation not found")
|
||||||
|
|
||||||
|
if conversation.user_id != user.id:
|
||||||
|
raise HTTPException(status_code=403, detail="Not authorized")
|
||||||
|
|
||||||
|
await service.delete(conversation_id)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{conversation_id}/messages", response_model=AddMessageResponse)
|
||||||
|
@logged()
|
||||||
|
async def add_message(
|
||||||
|
conversation_id: UUID,
|
||||||
|
request: AddMessageRequest,
|
||||||
|
session: AsyncSession = Depends(get_session),
|
||||||
|
user=Depends(require_auth),
|
||||||
|
) -> AddMessageResponse:
|
||||||
|
"""
|
||||||
|
Add a message to a conversation and get agent response.
|
||||||
|
|
||||||
|
This is the main endpoint for continuing conversations.
|
||||||
|
It:
|
||||||
|
1. Adds the user message
|
||||||
|
2. Checks if summarization is needed
|
||||||
|
3. Builds context from conversation history
|
||||||
|
4. Gets agent response
|
||||||
|
5. Adds agent response to conversation
|
||||||
|
6. Returns both messages
|
||||||
|
"""
|
||||||
|
service = ConversationService(session)
|
||||||
|
|
||||||
|
# Verify conversation exists and user owns it
|
||||||
|
conversation = await service.get(conversation_id)
|
||||||
|
if not conversation:
|
||||||
|
raise HTTPException(status_code=404, detail="Conversation not found")
|
||||||
|
|
||||||
|
if conversation.user_id != user.id:
|
||||||
|
raise HTTPException(status_code=403, detail="Not authorized")
|
||||||
|
|
||||||
|
# Add user message
|
||||||
|
user_message = await service.add_message(
|
||||||
|
conversation_id=conversation_id,
|
||||||
|
role="user",
|
||||||
|
content=request.content,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check if summarization needed before getting response
|
||||||
|
summarized = await service.summarize_if_needed(conversation_id)
|
||||||
|
|
||||||
|
# Get agent response with context
|
||||||
|
try:
|
||||||
|
response_text = await service.get_agent_response(
|
||||||
|
conversation_id=conversation_id,
|
||||||
|
user_message=request.content,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"Agent response failed: {e}")
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=500,
|
||||||
|
detail=f"Agent error: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Add assistant message
|
||||||
|
assistant_message = await service.add_message(
|
||||||
|
conversation_id=conversation_id,
|
||||||
|
role="assistant",
|
||||||
|
content=response_text,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Get updated conversation for total tokens
|
||||||
|
conversation = await service.get(conversation_id)
|
||||||
|
|
||||||
|
return AddMessageResponse(
|
||||||
|
user_message=MessageResponse.model_validate(user_message),
|
||||||
|
assistant_message=MessageResponse.model_validate(assistant_message),
|
||||||
|
total_tokens=conversation.total_tokens if conversation else 0,
|
||||||
|
summarized=summarized,
|
||||||
|
)
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
"""
|
||||||
|
Pydantic schemas for conversation API.
|
||||||
|
"""
|
||||||
|
from datetime import datetime
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
||||||
|
# === Request Schemas ===
|
||||||
|
|
||||||
|
class CreateConversationRequest(BaseModel):
|
||||||
|
"""Request to create a new conversation."""
|
||||||
|
agent_type: str = Field(default="explore", description="Agent type to use")
|
||||||
|
working_dir: str = Field(default=".", description="Working directory for agent")
|
||||||
|
title: str | None = Field(default=None, description="Optional conversation title")
|
||||||
|
|
||||||
|
|
||||||
|
class AddMessageRequest(BaseModel):
|
||||||
|
"""Request to add a message to a conversation."""
|
||||||
|
content: str = Field(..., min_length=1, description="Message content")
|
||||||
|
|
||||||
|
|
||||||
|
# === Response Schemas ===
|
||||||
|
|
||||||
|
class MessageResponse(BaseModel):
|
||||||
|
"""Response for a single message."""
|
||||||
|
id: UUID
|
||||||
|
role: str
|
||||||
|
content: str
|
||||||
|
token_count: int
|
||||||
|
is_summary: bool
|
||||||
|
created_at: datetime
|
||||||
|
|
||||||
|
model_config = {"from_attributes": True}
|
||||||
|
|
||||||
|
|
||||||
|
class ConversationResponse(BaseModel):
|
||||||
|
"""Response for conversation metadata."""
|
||||||
|
id: UUID
|
||||||
|
agent_type: str
|
||||||
|
title: str | None
|
||||||
|
working_dir: str
|
||||||
|
total_tokens: int
|
||||||
|
created_at: datetime
|
||||||
|
updated_at: datetime | None
|
||||||
|
|
||||||
|
model_config = {"from_attributes": True}
|
||||||
|
|
||||||
|
|
||||||
|
class ConversationDetailResponse(BaseModel):
|
||||||
|
"""Response for conversation with messages."""
|
||||||
|
id: UUID
|
||||||
|
agent_type: str
|
||||||
|
title: str | None
|
||||||
|
working_dir: str
|
||||||
|
total_tokens: int
|
||||||
|
created_at: datetime
|
||||||
|
updated_at: datetime | None
|
||||||
|
messages: list[MessageResponse]
|
||||||
|
|
||||||
|
model_config = {"from_attributes": True}
|
||||||
|
|
||||||
|
|
||||||
|
class ConversationListResponse(BaseModel):
|
||||||
|
"""Response for listing conversations."""
|
||||||
|
conversations: list[ConversationResponse]
|
||||||
|
total: int
|
||||||
|
|
||||||
|
|
||||||
|
class AddMessageResponse(BaseModel):
|
||||||
|
"""Response after adding a message (includes agent response)."""
|
||||||
|
user_message: MessageResponse
|
||||||
|
assistant_message: MessageResponse
|
||||||
|
total_tokens: int
|
||||||
|
summarized: bool = Field(
|
||||||
|
default=False,
|
||||||
|
description="Whether context was summarized due to token limit"
|
||||||
|
)
|
||||||
@@ -0,0 +1,354 @@
|
|||||||
|
"""
|
||||||
|
Conversation service - Business logic for conversation management.
|
||||||
|
|
||||||
|
Handles CRUD operations, context building, and summarization triggers.
|
||||||
|
"""
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
from sqlalchemy import select, func
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
from sqlalchemy.orm import selectinload
|
||||||
|
|
||||||
|
from src.domains.agents.base import get_agent
|
||||||
|
from src.domains.conversations.models import Conversation, Message
|
||||||
|
from src.domains.conversations.summarize import generate_summary
|
||||||
|
from src.shared.config import get_settings
|
||||||
|
from src.shared.logging import get_logger
|
||||||
|
from src.shared.tokens import count_tokens
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class ConversationService:
|
||||||
|
"""
|
||||||
|
Service for managing conversations and messages.
|
||||||
|
|
||||||
|
Handles:
|
||||||
|
- CRUD operations for conversations and messages
|
||||||
|
- Context building for agent prompts
|
||||||
|
- Automatic summarization when approaching token limits
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, session: AsyncSession):
|
||||||
|
"""
|
||||||
|
Initialize with database session.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Async SQLAlchemy session
|
||||||
|
"""
|
||||||
|
self.session = session
|
||||||
|
self.settings = get_settings()
|
||||||
|
|
||||||
|
# === Conversation CRUD ===
|
||||||
|
|
||||||
|
async def create(
|
||||||
|
self,
|
||||||
|
user_id: str,
|
||||||
|
agent_type: str = "explore",
|
||||||
|
working_dir: str = ".",
|
||||||
|
title: str | None = None,
|
||||||
|
) -> Conversation:
|
||||||
|
"""
|
||||||
|
Create a new conversation.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
user_id: Owner's user ID
|
||||||
|
agent_type: Type of agent for this conversation
|
||||||
|
working_dir: Working directory for agent
|
||||||
|
title: Optional title (auto-generated from first message if None)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Created Conversation object
|
||||||
|
"""
|
||||||
|
conversation = Conversation(
|
||||||
|
user_id=user_id,
|
||||||
|
agent_type=agent_type,
|
||||||
|
working_dir=working_dir,
|
||||||
|
title=title,
|
||||||
|
)
|
||||||
|
self.session.add(conversation)
|
||||||
|
await self.session.flush()
|
||||||
|
logger.info(f"Created conversation {conversation.id} for user {user_id}")
|
||||||
|
return conversation
|
||||||
|
|
||||||
|
async def get(self, conversation_id: UUID) -> Conversation | None:
|
||||||
|
"""Get conversation by ID without messages."""
|
||||||
|
result = await self.session.execute(
|
||||||
|
select(Conversation).where(Conversation.id == conversation_id)
|
||||||
|
)
|
||||||
|
return result.scalar_one_or_none()
|
||||||
|
|
||||||
|
async def get_with_messages(self, conversation_id: UUID) -> Conversation | None:
|
||||||
|
"""Get conversation by ID with messages loaded."""
|
||||||
|
result = await self.session.execute(
|
||||||
|
select(Conversation)
|
||||||
|
.options(selectinload(Conversation.messages))
|
||||||
|
.where(Conversation.id == conversation_id)
|
||||||
|
)
|
||||||
|
return result.scalar_one_or_none()
|
||||||
|
|
||||||
|
async def list_by_user(
|
||||||
|
self,
|
||||||
|
user_id: str,
|
||||||
|
limit: int = 50,
|
||||||
|
offset: int = 0,
|
||||||
|
) -> tuple[list[Conversation], int]:
|
||||||
|
"""
|
||||||
|
List conversations for a user.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
user_id: User ID to filter by
|
||||||
|
limit: Maximum results to return
|
||||||
|
offset: Offset for pagination
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (conversations, total_count)
|
||||||
|
"""
|
||||||
|
# Get total count
|
||||||
|
count_result = await self.session.execute(
|
||||||
|
select(func.count(Conversation.id))
|
||||||
|
.where(Conversation.user_id == user_id)
|
||||||
|
)
|
||||||
|
total = count_result.scalar() or 0
|
||||||
|
|
||||||
|
# Get conversations
|
||||||
|
result = await self.session.execute(
|
||||||
|
select(Conversation)
|
||||||
|
.where(Conversation.user_id == user_id)
|
||||||
|
.order_by(Conversation.updated_at.desc())
|
||||||
|
.limit(limit)
|
||||||
|
.offset(offset)
|
||||||
|
)
|
||||||
|
conversations = list(result.scalars().all())
|
||||||
|
|
||||||
|
return conversations, total
|
||||||
|
|
||||||
|
async def delete(self, conversation_id: UUID) -> bool:
|
||||||
|
"""Delete a conversation and all its messages."""
|
||||||
|
conversation = await self.get(conversation_id)
|
||||||
|
if conversation:
|
||||||
|
await self.session.delete(conversation)
|
||||||
|
logger.info(f"Deleted conversation {conversation_id}")
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
# === Message Operations ===
|
||||||
|
|
||||||
|
async def add_message(
|
||||||
|
self,
|
||||||
|
conversation_id: UUID,
|
||||||
|
role: str,
|
||||||
|
content: str,
|
||||||
|
) -> Message:
|
||||||
|
"""
|
||||||
|
Add a message to a conversation.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
conversation_id: Conversation to add to
|
||||||
|
role: Message role (user, assistant, system, summary)
|
||||||
|
content: Message content
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Created Message object
|
||||||
|
"""
|
||||||
|
# Count tokens
|
||||||
|
token_count = count_tokens(content)
|
||||||
|
|
||||||
|
message = Message(
|
||||||
|
conversation_id=conversation_id,
|
||||||
|
role=role,
|
||||||
|
content=content,
|
||||||
|
token_count=token_count,
|
||||||
|
)
|
||||||
|
self.session.add(message)
|
||||||
|
|
||||||
|
# Update conversation total tokens
|
||||||
|
conversation = await self.get(conversation_id)
|
||||||
|
if conversation:
|
||||||
|
conversation.total_tokens += token_count
|
||||||
|
|
||||||
|
# Auto-generate title from first user message
|
||||||
|
if conversation.title is None and role == "user":
|
||||||
|
conversation.title = content[:100] + ("..." if len(content) > 100 else "")
|
||||||
|
|
||||||
|
await self.session.flush()
|
||||||
|
return message
|
||||||
|
|
||||||
|
# === Context Building ===
|
||||||
|
|
||||||
|
def build_context_prompt(
|
||||||
|
self,
|
||||||
|
messages: list[Message],
|
||||||
|
current_message: str,
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Build a prompt with conversation context.
|
||||||
|
|
||||||
|
Includes summary (if exists) and recent messages.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: All conversation messages
|
||||||
|
current_message: The current user message
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Formatted prompt with context
|
||||||
|
"""
|
||||||
|
parts = []
|
||||||
|
|
||||||
|
# Find most recent summary
|
||||||
|
summaries = [m for m in messages if m.is_summary]
|
||||||
|
if summaries:
|
||||||
|
latest_summary = summaries[-1]
|
||||||
|
parts.append(
|
||||||
|
f"<conversation_summary>\n{latest_summary.content}\n</conversation_summary>"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Get recent non-summary messages
|
||||||
|
recent = [m for m in messages if not m.is_summary]
|
||||||
|
keep_count = self.settings.keep_recent_messages
|
||||||
|
recent = recent[-keep_count:] if len(recent) > keep_count else recent
|
||||||
|
|
||||||
|
if recent:
|
||||||
|
parts.append("<recent_conversation>")
|
||||||
|
for msg in recent:
|
||||||
|
role_label = msg.role.upper()
|
||||||
|
parts.append(f"{role_label}: {msg.content}")
|
||||||
|
parts.append("</recent_conversation>")
|
||||||
|
|
||||||
|
# Add current message
|
||||||
|
parts.append(f"<current_request>\n{current_message}\n</current_request>")
|
||||||
|
|
||||||
|
return "\n\n".join(parts)
|
||||||
|
|
||||||
|
# === Agent Integration ===
|
||||||
|
|
||||||
|
async def get_agent_response(
|
||||||
|
self,
|
||||||
|
conversation_id: UUID,
|
||||||
|
user_message: str,
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Get agent response with conversation context.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
conversation_id: Conversation ID
|
||||||
|
user_message: Current user message
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Agent's response text
|
||||||
|
"""
|
||||||
|
conversation = await self.get_with_messages(conversation_id)
|
||||||
|
if not conversation:
|
||||||
|
raise ValueError(f"Conversation {conversation_id} not found")
|
||||||
|
|
||||||
|
agent = get_agent(conversation.agent_type)
|
||||||
|
if not agent:
|
||||||
|
raise ValueError(f"Unknown agent type: {conversation.agent_type}")
|
||||||
|
|
||||||
|
# Build context prompt
|
||||||
|
context_prompt = self.build_context_prompt(
|
||||||
|
conversation.messages,
|
||||||
|
user_message,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Run agent
|
||||||
|
response = await agent.run(
|
||||||
|
context_prompt,
|
||||||
|
working_dir=conversation.working_dir,
|
||||||
|
)
|
||||||
|
|
||||||
|
return response
|
||||||
|
|
||||||
|
# === Summarization ===
|
||||||
|
|
||||||
|
async def should_summarize(self, conversation_id: UUID) -> bool:
|
||||||
|
"""
|
||||||
|
Check if conversation needs summarization.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
conversation_id: Conversation to check
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if summarization should be triggered
|
||||||
|
"""
|
||||||
|
conversation = await self.get(conversation_id)
|
||||||
|
if not conversation:
|
||||||
|
return False
|
||||||
|
|
||||||
|
threshold = self.settings.max_context_tokens * self.settings.summarization_threshold
|
||||||
|
return conversation.total_tokens > threshold
|
||||||
|
|
||||||
|
async def summarize_if_needed(self, conversation_id: UUID) -> bool:
|
||||||
|
"""
|
||||||
|
Summarize old messages if approaching token limit.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
conversation_id: Conversation to check and potentially summarize
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if summarization was performed
|
||||||
|
"""
|
||||||
|
if not await self.should_summarize(conversation_id):
|
||||||
|
return False
|
||||||
|
|
||||||
|
conversation = await self.get_with_messages(conversation_id)
|
||||||
|
if not conversation:
|
||||||
|
return False
|
||||||
|
|
||||||
|
messages = conversation.messages
|
||||||
|
keep_count = self.settings.keep_recent_messages
|
||||||
|
|
||||||
|
# Don't summarize if not enough messages
|
||||||
|
if len(messages) <= keep_count + 1:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Get messages to summarize (exclude recent and existing summaries)
|
||||||
|
non_summary_msgs = [m for m in messages if not m.is_summary]
|
||||||
|
to_summarize = non_summary_msgs[:-keep_count]
|
||||||
|
|
||||||
|
if not to_summarize:
|
||||||
|
return False
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"Summarizing {len(to_summarize)} messages in conversation {conversation_id}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Generate summary
|
||||||
|
summary_text = await generate_summary(
|
||||||
|
to_summarize,
|
||||||
|
working_dir=conversation.working_dir,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Get ID of last summarized message
|
||||||
|
last_summarized_id = to_summarize[-1].id
|
||||||
|
|
||||||
|
# Calculate tokens being removed
|
||||||
|
removed_tokens = sum(m.token_count for m in to_summarize)
|
||||||
|
summary_tokens = count_tokens(summary_text)
|
||||||
|
|
||||||
|
# Add summary message
|
||||||
|
summary_message = Message(
|
||||||
|
conversation_id=conversation_id,
|
||||||
|
role="summary",
|
||||||
|
content=summary_text,
|
||||||
|
token_count=summary_tokens,
|
||||||
|
is_summary=True,
|
||||||
|
summarizes_up_to=last_summarized_id,
|
||||||
|
)
|
||||||
|
self.session.add(summary_message)
|
||||||
|
|
||||||
|
# Mark old messages as summarized (soft delete by excluding from context)
|
||||||
|
for msg in to_summarize:
|
||||||
|
msg.is_summary = True # Reuse flag to mark as "summarized away"
|
||||||
|
|
||||||
|
# Update conversation token count
|
||||||
|
conversation.total_tokens = conversation.total_tokens - removed_tokens + summary_tokens
|
||||||
|
|
||||||
|
await self.session.flush()
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"Summarization complete: removed {removed_tokens} tokens, "
|
||||||
|
f"added {summary_tokens} token summary"
|
||||||
|
)
|
||||||
|
|
||||||
|
return True
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
"""
|
||||||
|
Context summarization for conversations.
|
||||||
|
|
||||||
|
Compresses old messages when approaching token limits.
|
||||||
|
"""
|
||||||
|
from src.domains.conversations.models import Message
|
||||||
|
from src.shared.logging import get_logger
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
SUMMARIZE_PROMPT = """Summarize this conversation history concisely for context preservation.
|
||||||
|
|
||||||
|
Focus on:
|
||||||
|
- Key decisions made and their rationale
|
||||||
|
- Important files, functions, or code discussed
|
||||||
|
- Current task state and progress
|
||||||
|
- Any unresolved questions or blockers
|
||||||
|
- Technical details that would be needed to continue the work
|
||||||
|
|
||||||
|
Keep the summary under 500 words. Be factual and technical, not conversational.
|
||||||
|
Preserve specific file paths, function names, and code references.
|
||||||
|
|
||||||
|
CONVERSATION HISTORY:
|
||||||
|
{history}
|
||||||
|
|
||||||
|
CONCISE SUMMARY:"""
|
||||||
|
|
||||||
|
|
||||||
|
def format_messages_for_summary(messages: list[Message]) -> str:
|
||||||
|
"""
|
||||||
|
Format messages into a string for summarization.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: List of Message objects to format
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Formatted conversation string
|
||||||
|
"""
|
||||||
|
parts = []
|
||||||
|
for msg in messages:
|
||||||
|
if msg.is_summary:
|
||||||
|
parts.append(f"[Previous Summary]: {msg.content}")
|
||||||
|
else:
|
||||||
|
role = msg.role.upper()
|
||||||
|
parts.append(f"{role}: {msg.content}")
|
||||||
|
return "\n\n".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
async def generate_summary(
|
||||||
|
messages: list[Message],
|
||||||
|
working_dir: str = "."
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Generate a summary of conversation messages using the Explore agent.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: Messages to summarize
|
||||||
|
working_dir: Working directory for agent context
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Summary text
|
||||||
|
"""
|
||||||
|
from src.domains.agents.explore import explore
|
||||||
|
|
||||||
|
history = format_messages_for_summary(messages)
|
||||||
|
prompt = SUMMARIZE_PROMPT.format(history=history)
|
||||||
|
|
||||||
|
logger.info(f"Generating summary for {len(messages)} messages")
|
||||||
|
|
||||||
|
try:
|
||||||
|
summary = await explore(prompt, working_dir=working_dir)
|
||||||
|
return summary.strip()
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Summary generation failed: {e}")
|
||||||
|
# Fallback: create a simple truncated summary
|
||||||
|
return _fallback_summary(messages)
|
||||||
|
|
||||||
|
|
||||||
|
def _fallback_summary(messages: list[Message]) -> str:
|
||||||
|
"""
|
||||||
|
Create a simple fallback summary if agent summarization fails.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: Messages to summarize
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Basic summary string
|
||||||
|
"""
|
||||||
|
# Take first and last few messages
|
||||||
|
if len(messages) <= 4:
|
||||||
|
return format_messages_for_summary(messages)
|
||||||
|
|
||||||
|
first_two = messages[:2]
|
||||||
|
last_two = messages[-2:]
|
||||||
|
|
||||||
|
parts = [
|
||||||
|
"Conversation started with:",
|
||||||
|
format_messages_for_summary(first_two),
|
||||||
|
f"\n[... {len(messages) - 4} messages omitted ...]\n",
|
||||||
|
"Most recent exchange:",
|
||||||
|
format_messages_for_summary(last_two),
|
||||||
|
]
|
||||||
|
return "\n".join(parts)
|
||||||
@@ -8,6 +8,7 @@ from fastapi import APIRouter
|
|||||||
|
|
||||||
from src.domains.health.router import router as health_router
|
from src.domains.health.router import router as health_router
|
||||||
from src.domains.agents.router import router as agents_router
|
from src.domains.agents.router import router as agents_router
|
||||||
|
from src.domains.conversations.router import router as conversations_router
|
||||||
|
|
||||||
# from src.domains.auth.router import router as auth_router
|
# from src.domains.auth.router import router as auth_router
|
||||||
# from src.domains.tools.router import router as tools_router
|
# from src.domains.tools.router import router as tools_router
|
||||||
@@ -20,6 +21,9 @@ root_router.include_router(health_router)
|
|||||||
# Agents domain (prefix defined in router)
|
# Agents domain (prefix defined in router)
|
||||||
root_router.include_router(agents_router)
|
root_router.include_router(agents_router)
|
||||||
|
|
||||||
|
# Conversations domain (prefix defined in router)
|
||||||
|
root_router.include_router(conversations_router)
|
||||||
|
|
||||||
# Auth domain
|
# Auth domain
|
||||||
# root_router.include_router(auth_router, prefix="/auth", tags=["Auth"])
|
# root_router.include_router(auth_router, prefix="/auth", tags=["Auth"])
|
||||||
|
|
||||||
|
|||||||
@@ -31,13 +31,18 @@ async def lifespan(app: FastAPI):
|
|||||||
logger.info(f"Port: {settings.port}")
|
logger.info(f"Port: {settings.port}")
|
||||||
logger.info(f"Ollama: {settings.ollama_url}")
|
logger.info(f"Ollama: {settings.ollama_url}")
|
||||||
logger.info(f"Agent model: {settings.ollama_agent_model}")
|
logger.info(f"Agent model: {settings.ollama_agent_model}")
|
||||||
|
logger.info(f"Database: {settings.database_url}")
|
||||||
logger.info("=" * 60)
|
logger.info("=" * 60)
|
||||||
|
|
||||||
# TODO: Initialize resources (LLM clients, etc.)
|
|
||||||
|
|
||||||
yield
|
yield
|
||||||
|
|
||||||
# Cleanup
|
# Cleanup
|
||||||
|
from src.db import get_database
|
||||||
|
try:
|
||||||
|
database = get_database()
|
||||||
|
await database.close()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
logger.info("Shutting down")
|
logger.info("Shutting down")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -80,9 +80,15 @@ class Settings(BaseSettings):
|
|||||||
sandbox_enabled: bool = True
|
sandbox_enabled: bool = True
|
||||||
allowed_paths: list[str] | None = None
|
allowed_paths: list[str] | None = None
|
||||||
|
|
||||||
# Sessions
|
# Database
|
||||||
|
database_url: str = "sqlite+aiosqlite:///./webber.db"
|
||||||
|
|
||||||
|
# Sessions & Context
|
||||||
session_ttl_hours: int = 24
|
session_ttl_hours: int = 24
|
||||||
max_context_tokens: int = 128000
|
max_context_tokens: int = 128000
|
||||||
|
summarization_threshold: float = 0.8 # Summarize at 80% of max tokens
|
||||||
|
summarization_target_tokens: int = 500 # Target summary size
|
||||||
|
keep_recent_messages: int = 6 # Messages to keep unsummarized (3 turns)
|
||||||
|
|
||||||
model_config = SettingsConfigDict(
|
model_config = SettingsConfigDict(
|
||||||
env_file=".env",
|
env_file=".env",
|
||||||
|
|||||||
@@ -0,0 +1,74 @@
|
|||||||
|
"""
|
||||||
|
Token counting utilities for context management.
|
||||||
|
|
||||||
|
Uses litellm for accurate multi-model token counting.
|
||||||
|
"""
|
||||||
|
from src.shared.logging import get_logger
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
# Default model for token counting (Mistral Nemo)
|
||||||
|
DEFAULT_MODEL = "mistral/mistral-nemo"
|
||||||
|
|
||||||
|
|
||||||
|
def count_tokens(text: str, model: str = DEFAULT_MODEL) -> int:
|
||||||
|
"""
|
||||||
|
Count tokens in a text string.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: Text to count tokens for
|
||||||
|
model: Model identifier for tokenizer selection
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Token count
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from litellm import token_counter
|
||||||
|
return token_counter(model=model, text=text)
|
||||||
|
except Exception as e:
|
||||||
|
# Fallback to rough estimate if litellm fails
|
||||||
|
logger.warning(f"Token counting failed, using estimate: {e}")
|
||||||
|
return len(text) // 4
|
||||||
|
|
||||||
|
|
||||||
|
def count_message_tokens(
|
||||||
|
messages: list[dict[str, str]],
|
||||||
|
model: str = DEFAULT_MODEL
|
||||||
|
) -> int:
|
||||||
|
"""
|
||||||
|
Count tokens for a list of chat messages.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: List of message dicts with 'role' and 'content' keys
|
||||||
|
model: Model identifier for tokenizer selection
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Total token count including message overhead
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from litellm import token_counter
|
||||||
|
return token_counter(model=model, messages=messages)
|
||||||
|
except Exception as e:
|
||||||
|
# Fallback to rough estimate
|
||||||
|
logger.warning(f"Token counting failed, using estimate: {e}")
|
||||||
|
total = 0
|
||||||
|
for msg in messages:
|
||||||
|
total += len(msg.get("content", "")) // 4
|
||||||
|
total += 4 # Overhead per message
|
||||||
|
return total
|
||||||
|
|
||||||
|
|
||||||
|
def estimate_tokens(text: str) -> int:
|
||||||
|
"""
|
||||||
|
Quick token estimate without external library.
|
||||||
|
|
||||||
|
Uses ~4 characters per token heuristic.
|
||||||
|
Less accurate but faster for rough estimates.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: Text to estimate
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Estimated token count
|
||||||
|
"""
|
||||||
|
return len(text) // 4
|
||||||
@@ -0,0 +1,299 @@
|
|||||||
|
"""
|
||||||
|
Tests for conversations domain.
|
||||||
|
|
||||||
|
Tests conversation CRUD, context building, and API endpoints.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
from src.domains.conversations.models import Conversation, Message
|
||||||
|
from src.domains.conversations.schemas import (
|
||||||
|
CreateConversationRequest,
|
||||||
|
AddMessageRequest,
|
||||||
|
ConversationResponse,
|
||||||
|
MessageResponse,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestConversationModels:
|
||||||
|
"""Tests for conversation database models."""
|
||||||
|
|
||||||
|
def test_conversation_creation(self):
|
||||||
|
"""Test Conversation model creation with explicit values."""
|
||||||
|
conv = Conversation(
|
||||||
|
user_id="test-user",
|
||||||
|
agent_type="explore",
|
||||||
|
working_dir=".",
|
||||||
|
total_tokens=0,
|
||||||
|
)
|
||||||
|
assert conv.user_id == "test-user"
|
||||||
|
assert conv.agent_type == "explore"
|
||||||
|
assert conv.working_dir == "."
|
||||||
|
assert conv.total_tokens == 0
|
||||||
|
|
||||||
|
def test_conversation_with_values(self):
|
||||||
|
"""Test Conversation with explicit values."""
|
||||||
|
conv = Conversation(
|
||||||
|
user_id="test-user",
|
||||||
|
agent_type="plan",
|
||||||
|
working_dir="/tmp/project",
|
||||||
|
title="Test Conversation",
|
||||||
|
)
|
||||||
|
assert conv.agent_type == "plan"
|
||||||
|
assert conv.working_dir == "/tmp/project"
|
||||||
|
assert conv.title == "Test Conversation"
|
||||||
|
|
||||||
|
def test_message_creation(self):
|
||||||
|
"""Test Message model creation with explicit values."""
|
||||||
|
msg = Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="user",
|
||||||
|
content="Hello",
|
||||||
|
token_count=0,
|
||||||
|
is_summary=False,
|
||||||
|
)
|
||||||
|
assert msg.role == "user"
|
||||||
|
assert msg.content == "Hello"
|
||||||
|
assert msg.token_count == 0
|
||||||
|
assert msg.is_summary is False
|
||||||
|
|
||||||
|
def test_message_repr(self):
|
||||||
|
"""Test Message string representation."""
|
||||||
|
msg = Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="user",
|
||||||
|
content="This is a test message",
|
||||||
|
)
|
||||||
|
repr_str = repr(msg)
|
||||||
|
assert "user" in repr_str
|
||||||
|
assert "This is a test" in repr_str
|
||||||
|
|
||||||
|
|
||||||
|
class TestConversationSchemas:
|
||||||
|
"""Tests for Pydantic schemas."""
|
||||||
|
|
||||||
|
def test_create_request_defaults(self):
|
||||||
|
"""Test CreateConversationRequest defaults."""
|
||||||
|
request = CreateConversationRequest()
|
||||||
|
assert request.agent_type == "explore"
|
||||||
|
assert request.working_dir == "."
|
||||||
|
assert request.title is None
|
||||||
|
|
||||||
|
def test_create_request_custom(self):
|
||||||
|
"""Test CreateConversationRequest with values."""
|
||||||
|
request = CreateConversationRequest(
|
||||||
|
agent_type="task",
|
||||||
|
working_dir="/home/user/project",
|
||||||
|
title="My Task",
|
||||||
|
)
|
||||||
|
assert request.agent_type == "task"
|
||||||
|
assert request.working_dir == "/home/user/project"
|
||||||
|
assert request.title == "My Task"
|
||||||
|
|
||||||
|
def test_add_message_request_valid(self):
|
||||||
|
"""Test AddMessageRequest validation."""
|
||||||
|
request = AddMessageRequest(content="Hello, world!")
|
||||||
|
assert request.content == "Hello, world!"
|
||||||
|
|
||||||
|
def test_add_message_request_empty_fails(self):
|
||||||
|
"""Test that empty content fails validation."""
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
AddMessageRequest(content="")
|
||||||
|
|
||||||
|
|
||||||
|
class TestConversationAPI:
|
||||||
|
"""Tests for conversation API endpoints."""
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_create_conversation(self, auth_client):
|
||||||
|
"""Test creating a conversation."""
|
||||||
|
response = await auth_client.post(
|
||||||
|
"/conversations/",
|
||||||
|
json={"agent_type": "explore", "working_dir": "."}
|
||||||
|
)
|
||||||
|
assert response.status_code == 201
|
||||||
|
data = response.json()
|
||||||
|
assert "id" in data
|
||||||
|
assert data["agent_type"] == "explore"
|
||||||
|
assert data["total_tokens"] == 0
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_create_conversation_with_title(self, auth_client):
|
||||||
|
"""Test creating a conversation with title."""
|
||||||
|
response = await auth_client.post(
|
||||||
|
"/conversations/",
|
||||||
|
json={
|
||||||
|
"agent_type": "plan",
|
||||||
|
"working_dir": "/tmp",
|
||||||
|
"title": "Planning Session"
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assert response.status_code == 201
|
||||||
|
data = response.json()
|
||||||
|
assert data["title"] == "Planning Session"
|
||||||
|
assert data["agent_type"] == "plan"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_list_conversations_empty(self, auth_client):
|
||||||
|
"""Test listing conversations when empty."""
|
||||||
|
response = await auth_client.get("/conversations/")
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert "conversations" in data
|
||||||
|
assert "total" in data
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_conversation_not_found(self, auth_client):
|
||||||
|
"""Test getting non-existent conversation."""
|
||||||
|
fake_id = uuid4()
|
||||||
|
response = await auth_client.get(f"/conversations/{fake_id}")
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_delete_conversation_not_found(self, auth_client):
|
||||||
|
"""Test deleting non-existent conversation."""
|
||||||
|
fake_id = uuid4()
|
||||||
|
response = await auth_client.delete(f"/conversations/{fake_id}")
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_add_message_not_found(self, auth_client):
|
||||||
|
"""Test adding message to non-existent conversation."""
|
||||||
|
fake_id = uuid4()
|
||||||
|
response = await auth_client.post(
|
||||||
|
f"/conversations/{fake_id}/messages",
|
||||||
|
json={"content": "Hello"}
|
||||||
|
)
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
class TestConversationService:
|
||||||
|
"""Tests for ConversationService business logic."""
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_context_prompt_no_history(self):
|
||||||
|
"""Test building context prompt with no history."""
|
||||||
|
from src.domains.conversations.service import ConversationService
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
# Create mock session
|
||||||
|
mock_session = MagicMock()
|
||||||
|
service = ConversationService(mock_session)
|
||||||
|
|
||||||
|
prompt = service.build_context_prompt([], "What files are here?")
|
||||||
|
|
||||||
|
assert "<current_request>" in prompt
|
||||||
|
assert "What files are here?" in prompt
|
||||||
|
assert "<recent_conversation>" not in prompt
|
||||||
|
assert "<conversation_summary>" not in prompt
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_context_prompt_with_history(self):
|
||||||
|
"""Test building context prompt with message history."""
|
||||||
|
from src.domains.conversations.service import ConversationService
|
||||||
|
from src.domains.conversations.models import Message
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
mock_session = MagicMock()
|
||||||
|
service = ConversationService(mock_session)
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="user",
|
||||||
|
content="Find Python files",
|
||||||
|
),
|
||||||
|
Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="assistant",
|
||||||
|
content="Found 10 Python files.",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
prompt = service.build_context_prompt(messages, "Show the largest")
|
||||||
|
|
||||||
|
assert "<recent_conversation>" in prompt
|
||||||
|
assert "USER: Find Python files" in prompt
|
||||||
|
assert "ASSISTANT: Found 10 Python files" in prompt
|
||||||
|
assert "<current_request>" in prompt
|
||||||
|
assert "Show the largest" in prompt
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_context_prompt_with_summary(self):
|
||||||
|
"""Test building context prompt with summary message."""
|
||||||
|
from src.domains.conversations.service import ConversationService
|
||||||
|
from src.domains.conversations.models import Message
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
mock_session = MagicMock()
|
||||||
|
service = ConversationService(mock_session)
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="summary",
|
||||||
|
content="Previously discussed: project setup",
|
||||||
|
is_summary=True,
|
||||||
|
),
|
||||||
|
Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="user",
|
||||||
|
content="Now what?",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
prompt = service.build_context_prompt(messages, "Continue")
|
||||||
|
|
||||||
|
assert "<conversation_summary>" in prompt
|
||||||
|
assert "Previously discussed: project setup" in prompt
|
||||||
|
|
||||||
|
|
||||||
|
class TestSummarization:
|
||||||
|
"""Tests for conversation summarization."""
|
||||||
|
|
||||||
|
def test_format_messages_for_summary(self):
|
||||||
|
"""Test formatting messages for summarization."""
|
||||||
|
from src.domains.conversations.summarize import format_messages_for_summary
|
||||||
|
from src.domains.conversations.models import Message
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="user",
|
||||||
|
content="Hello",
|
||||||
|
),
|
||||||
|
Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="assistant",
|
||||||
|
content="Hi there!",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
formatted = format_messages_for_summary(messages)
|
||||||
|
|
||||||
|
assert "USER: Hello" in formatted
|
||||||
|
assert "ASSISTANT: Hi there!" in formatted
|
||||||
|
|
||||||
|
def test_format_messages_with_summary(self):
|
||||||
|
"""Test formatting messages that include a summary."""
|
||||||
|
from src.domains.conversations.summarize import format_messages_for_summary
|
||||||
|
from src.domains.conversations.models import Message
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="summary",
|
||||||
|
content="Previous context summary",
|
||||||
|
is_summary=True,
|
||||||
|
),
|
||||||
|
Message(
|
||||||
|
conversation_id=uuid4(),
|
||||||
|
role="user",
|
||||||
|
content="Continue",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
formatted = format_messages_for_summary(messages)
|
||||||
|
|
||||||
|
assert "[Previous Summary]" in formatted
|
||||||
|
assert "Previous context summary" in formatted
|
||||||
@@ -0,0 +1,152 @@
|
|||||||
|
"""
|
||||||
|
Tests for the Plan agent.
|
||||||
|
|
||||||
|
Tests registration, API endpoints, and tool restrictions.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.domains.agents.base import get_agent, list_agents
|
||||||
|
from src.domains.agents.plan import plan_agent, PlanAgentImpl
|
||||||
|
|
||||||
|
|
||||||
|
class TestPlanAgentRegistration:
|
||||||
|
"""Tests for Plan agent registration."""
|
||||||
|
|
||||||
|
def test_plan_agent_registered(self):
|
||||||
|
"""Test that plan agent is registered in registry."""
|
||||||
|
agent = get_agent("plan")
|
||||||
|
assert agent is not None
|
||||||
|
assert agent.name == "plan"
|
||||||
|
|
||||||
|
def test_plan_agent_in_list(self):
|
||||||
|
"""Test that plan agent appears in agent list."""
|
||||||
|
agents = list_agents()
|
||||||
|
names = [a["name"] for a in agents]
|
||||||
|
assert "plan" in names
|
||||||
|
|
||||||
|
def test_plan_agent_has_description(self):
|
||||||
|
"""Test that plan agent has a description."""
|
||||||
|
agent = get_agent("plan")
|
||||||
|
assert agent is not None
|
||||||
|
assert len(agent.description) > 0
|
||||||
|
assert "plan" in agent.description.lower() or "architect" in agent.description.lower()
|
||||||
|
|
||||||
|
def test_plan_agent_singleton(self):
|
||||||
|
"""Test that plan_agent is the registered instance."""
|
||||||
|
registered = get_agent("plan")
|
||||||
|
assert registered is plan_agent
|
||||||
|
|
||||||
|
def test_plan_agent_is_correct_type(self):
|
||||||
|
"""Test that plan agent is correct implementation type."""
|
||||||
|
assert isinstance(plan_agent, PlanAgentImpl)
|
||||||
|
|
||||||
|
|
||||||
|
class TestPlanAgentTools:
|
||||||
|
"""Tests for Plan agent tool restrictions."""
|
||||||
|
|
||||||
|
def test_plan_agent_has_read_only_tools(self):
|
||||||
|
"""Test that plan agent has read-only tools."""
|
||||||
|
# Access the underlying PydanticAI agent to check tools
|
||||||
|
agent = plan_agent.agent
|
||||||
|
tool_names = list(agent._function_toolset.tools.keys())
|
||||||
|
|
||||||
|
# Should have read-only tools
|
||||||
|
assert "read_file" in tool_names
|
||||||
|
assert "glob_files" in tool_names
|
||||||
|
assert "grep_content" in tool_names
|
||||||
|
assert "bash_readonly" in tool_names
|
||||||
|
|
||||||
|
def test_plan_agent_no_write_tools(self):
|
||||||
|
"""Test that plan agent does NOT have write tools."""
|
||||||
|
agent = plan_agent.agent
|
||||||
|
tool_names = list(agent._function_toolset.tools.keys())
|
||||||
|
|
||||||
|
# Should NOT have write tools
|
||||||
|
assert "edit_file" not in tool_names
|
||||||
|
assert "write_file" not in tool_names
|
||||||
|
assert "bash" not in tool_names
|
||||||
|
assert "web_search" not in tool_names
|
||||||
|
|
||||||
|
def test_plan_agent_tool_count(self):
|
||||||
|
"""Test that plan agent has exactly 4 tools."""
|
||||||
|
agent = plan_agent.agent
|
||||||
|
tool_count = len(agent._function_toolset.tools)
|
||||||
|
assert tool_count == 4
|
||||||
|
|
||||||
|
|
||||||
|
class TestPlanAgentAPI:
|
||||||
|
"""Tests for Plan agent REST API."""
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_list_agents_includes_plan(self, auth_client):
|
||||||
|
"""Test that agent list includes plan agent."""
|
||||||
|
response = await auth_client.get("/agents/")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
names = [a["name"] for a in data["agents"]]
|
||||||
|
assert "plan" in names
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_plan_agent_info(self, auth_client):
|
||||||
|
"""Test getting plan agent info."""
|
||||||
|
response = await auth_client.get("/agents/plan")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert data["name"] == "plan"
|
||||||
|
assert "description" in data
|
||||||
|
assert len(data["description"]) > 0
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_run_plan_with_invalid_body(self, auth_client):
|
||||||
|
"""Test running plan agent with invalid request."""
|
||||||
|
response = await auth_client.post(
|
||||||
|
"/agents/run",
|
||||||
|
json={
|
||||||
|
"agent_type": "plan",
|
||||||
|
# Missing prompt
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 422
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_stream_plan_with_invalid_body(self, auth_client):
|
||||||
|
"""Test streaming plan agent with invalid request."""
|
||||||
|
response = await auth_client.post(
|
||||||
|
"/agents/stream",
|
||||||
|
json={
|
||||||
|
"agent_type": "plan",
|
||||||
|
# Missing prompt
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 422
|
||||||
|
|
||||||
|
|
||||||
|
class TestPlanAgentProperties:
|
||||||
|
"""Tests for Plan agent properties and configuration."""
|
||||||
|
|
||||||
|
def test_plan_agent_name(self):
|
||||||
|
"""Test plan agent name property."""
|
||||||
|
assert plan_agent.name == "plan"
|
||||||
|
|
||||||
|
def test_plan_agent_description_not_empty(self):
|
||||||
|
"""Test plan agent description is not empty."""
|
||||||
|
assert plan_agent.description
|
||||||
|
assert len(plan_agent.description) > 10
|
||||||
|
|
||||||
|
def test_plan_agent_creates_agent_lazily(self):
|
||||||
|
"""Test that PydanticAI agent is created lazily."""
|
||||||
|
# Create a fresh instance
|
||||||
|
fresh_agent = PlanAgentImpl()
|
||||||
|
|
||||||
|
# _agent should be None before first access
|
||||||
|
assert fresh_agent._agent is None
|
||||||
|
|
||||||
|
# Access the agent property
|
||||||
|
_ = fresh_agent.agent
|
||||||
|
|
||||||
|
# Now _agent should be set
|
||||||
|
assert fresh_agent._agent is not None
|
||||||
@@ -0,0 +1,257 @@
|
|||||||
|
"""
|
||||||
|
Tests for the Task agent.
|
||||||
|
|
||||||
|
Tests registration, API endpoints, tool access, and spawn_agent functionality.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
|
from src.domains.agents.base import get_agent, list_agents
|
||||||
|
from src.domains.agents.task import task_agent, TaskAgentImpl
|
||||||
|
|
||||||
|
|
||||||
|
class TestTaskAgentRegistration:
|
||||||
|
"""Tests for Task agent registration."""
|
||||||
|
|
||||||
|
def test_task_agent_registered(self):
|
||||||
|
"""Test that task agent is registered in registry."""
|
||||||
|
agent = get_agent("task")
|
||||||
|
assert agent is not None
|
||||||
|
assert agent.name == "task"
|
||||||
|
|
||||||
|
def test_task_agent_in_list(self):
|
||||||
|
"""Test that task agent appears in agent list."""
|
||||||
|
agents = list_agents()
|
||||||
|
names = [a["name"] for a in agents]
|
||||||
|
assert "task" in names
|
||||||
|
|
||||||
|
def test_task_agent_has_description(self):
|
||||||
|
"""Test that task agent has a description."""
|
||||||
|
agent = get_agent("task")
|
||||||
|
assert agent is not None
|
||||||
|
assert len(agent.description) > 0
|
||||||
|
assert "task" in agent.description.lower() or "autonomous" in agent.description.lower()
|
||||||
|
|
||||||
|
def test_task_agent_singleton(self):
|
||||||
|
"""Test that task_agent is the registered instance."""
|
||||||
|
registered = get_agent("task")
|
||||||
|
assert registered is task_agent
|
||||||
|
|
||||||
|
def test_task_agent_is_correct_type(self):
|
||||||
|
"""Test that task agent is correct implementation type."""
|
||||||
|
assert isinstance(task_agent, TaskAgentImpl)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTaskAgentTools:
|
||||||
|
"""Tests for Task agent tool access."""
|
||||||
|
|
||||||
|
def test_task_agent_has_all_tools(self):
|
||||||
|
"""Test that task agent has all 9 tools."""
|
||||||
|
agent = task_agent.agent
|
||||||
|
tool_names = list(agent._function_toolset.tools.keys())
|
||||||
|
|
||||||
|
# Should have 9 tools total
|
||||||
|
assert len(tool_names) == 9
|
||||||
|
|
||||||
|
def test_task_agent_has_read_only_tools(self):
|
||||||
|
"""Test that task agent has read-only tools."""
|
||||||
|
agent = task_agent.agent
|
||||||
|
tool_names = list(agent._function_toolset.tools.keys())
|
||||||
|
|
||||||
|
assert "read_file" in tool_names
|
||||||
|
assert "glob_files" in tool_names
|
||||||
|
assert "grep_content" in tool_names
|
||||||
|
assert "bash_readonly" in tool_names
|
||||||
|
|
||||||
|
def test_task_agent_has_write_tools(self):
|
||||||
|
"""Test that task agent has write tools."""
|
||||||
|
agent = task_agent.agent
|
||||||
|
tool_names = list(agent._function_toolset.tools.keys())
|
||||||
|
|
||||||
|
assert "edit_file" in tool_names
|
||||||
|
assert "write_file" in tool_names
|
||||||
|
assert "bash" in tool_names
|
||||||
|
|
||||||
|
def test_task_agent_has_external_tools(self):
|
||||||
|
"""Test that task agent has external tools."""
|
||||||
|
agent = task_agent.agent
|
||||||
|
tool_names = list(agent._function_toolset.tools.keys())
|
||||||
|
|
||||||
|
assert "web_search" in tool_names
|
||||||
|
|
||||||
|
def test_task_agent_has_spawn_agent_tool(self):
|
||||||
|
"""Test that task agent has spawn_agent orchestration tool."""
|
||||||
|
agent = task_agent.agent
|
||||||
|
tool_names = list(agent._function_toolset.tools.keys())
|
||||||
|
|
||||||
|
assert "spawn_agent" in tool_names
|
||||||
|
|
||||||
|
|
||||||
|
class TestSpawnAgentTool:
|
||||||
|
"""Tests for spawn_agent orchestration functionality."""
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_spawn_explore_agent(self):
|
||||||
|
"""Test spawning an explore agent."""
|
||||||
|
from src.domains.agents.task.tools import register_task_tools
|
||||||
|
from src.domains.agents.base import AgentContext
|
||||||
|
from pydantic_ai import Agent, RunContext
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
# Create a mock context
|
||||||
|
ctx = MagicMock(spec=RunContext)
|
||||||
|
ctx.deps = AgentContext(
|
||||||
|
working_dir="/tmp",
|
||||||
|
allowed_paths=["/tmp"],
|
||||||
|
timeout_seconds=30
|
||||||
|
)
|
||||||
|
|
||||||
|
# Mock the explore agent
|
||||||
|
with patch("src.domains.agents.base.get_agent") as mock_get_agent:
|
||||||
|
mock_explore = AsyncMock()
|
||||||
|
mock_explore.run = AsyncMock(return_value="Found 5 Python files")
|
||||||
|
mock_get_agent.return_value = mock_explore
|
||||||
|
|
||||||
|
# Import and call spawn_agent directly
|
||||||
|
from src.domains.agents.task import tools
|
||||||
|
# We need to test the actual tool function
|
||||||
|
# For now, verify the explore agent would be called correctly
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_spawn_unknown_agent_returns_error(self):
|
||||||
|
"""Test that spawning unknown agent type returns error."""
|
||||||
|
from src.domains.agents.base import AgentContext
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
from pydantic_ai import RunContext
|
||||||
|
|
||||||
|
# We can't easily test the tool directly, but we can verify
|
||||||
|
# the agent type validation logic
|
||||||
|
allowed_types = ["explore", "plan"]
|
||||||
|
assert "nonexistent" not in allowed_types
|
||||||
|
assert "task" not in allowed_types # Task should be blocked
|
||||||
|
|
||||||
|
def test_spawn_task_agent_blocked(self):
|
||||||
|
"""Test that spawning nested task agents is blocked."""
|
||||||
|
# Verify the validation logic prevents recursion
|
||||||
|
# The spawn_agent tool should return an error for agent_type="task"
|
||||||
|
allowed_types = ["explore", "plan"]
|
||||||
|
assert "task" not in allowed_types
|
||||||
|
|
||||||
|
|
||||||
|
class TestTaskAgentAPI:
|
||||||
|
"""Tests for Task agent REST API."""
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_list_agents_includes_task(self, auth_client):
|
||||||
|
"""Test that agent list includes task agent."""
|
||||||
|
response = await auth_client.get("/agents/")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
names = [a["name"] for a in data["agents"]]
|
||||||
|
assert "task" in names
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_task_agent_info(self, auth_client):
|
||||||
|
"""Test getting task agent info."""
|
||||||
|
response = await auth_client.get("/agents/task")
|
||||||
|
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert data["name"] == "task"
|
||||||
|
assert "description" in data
|
||||||
|
assert len(data["description"]) > 0
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_run_task_with_invalid_body(self, auth_client):
|
||||||
|
"""Test running task agent with invalid request."""
|
||||||
|
response = await auth_client.post(
|
||||||
|
"/agents/run",
|
||||||
|
json={
|
||||||
|
"agent_type": "task",
|
||||||
|
# Missing prompt
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 422
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_stream_task_with_invalid_body(self, auth_client):
|
||||||
|
"""Test streaming task agent with invalid request."""
|
||||||
|
response = await auth_client.post(
|
||||||
|
"/agents/stream",
|
||||||
|
json={
|
||||||
|
"agent_type": "task",
|
||||||
|
# Missing prompt
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.status_code == 422
|
||||||
|
|
||||||
|
|
||||||
|
class TestTaskAgentProperties:
|
||||||
|
"""Tests for Task agent properties and configuration."""
|
||||||
|
|
||||||
|
def test_task_agent_name(self):
|
||||||
|
"""Test task agent name property."""
|
||||||
|
assert task_agent.name == "task"
|
||||||
|
|
||||||
|
def test_task_agent_description_not_empty(self):
|
||||||
|
"""Test task agent description is not empty."""
|
||||||
|
assert task_agent.description
|
||||||
|
assert len(task_agent.description) > 10
|
||||||
|
|
||||||
|
def test_task_agent_creates_agent_lazily(self):
|
||||||
|
"""Test that PydanticAI agent is created lazily."""
|
||||||
|
# Create a fresh instance
|
||||||
|
fresh_agent = TaskAgentImpl()
|
||||||
|
|
||||||
|
# _agent should be None before first access
|
||||||
|
assert fresh_agent._agent is None
|
||||||
|
|
||||||
|
# Access the agent property
|
||||||
|
_ = fresh_agent.agent
|
||||||
|
|
||||||
|
# Now _agent should be set
|
||||||
|
assert fresh_agent._agent is not None
|
||||||
|
|
||||||
|
|
||||||
|
class TestAllAgentsRegistered:
|
||||||
|
"""Tests to verify all three agents are registered."""
|
||||||
|
|
||||||
|
def test_all_agents_in_registry(self):
|
||||||
|
"""Test that explore, plan, and task agents are all registered."""
|
||||||
|
agents = list_agents()
|
||||||
|
names = [a["name"] for a in agents]
|
||||||
|
|
||||||
|
assert "explore" in names
|
||||||
|
assert "plan" in names
|
||||||
|
assert "task" in names
|
||||||
|
assert len(names) == 3
|
||||||
|
|
||||||
|
def test_agent_hierarchy(self):
|
||||||
|
"""Test the agent capability hierarchy."""
|
||||||
|
explore = get_agent("explore")
|
||||||
|
plan = get_agent("plan")
|
||||||
|
task = get_agent("task")
|
||||||
|
|
||||||
|
explore_tools = list(explore.agent._function_toolset.tools.keys())
|
||||||
|
plan_tools = list(plan.agent._function_toolset.tools.keys())
|
||||||
|
task_tools = list(task.agent._function_toolset.tools.keys())
|
||||||
|
|
||||||
|
# Explore has all tools (read + write)
|
||||||
|
assert "edit_file" in explore_tools
|
||||||
|
assert "write_file" in explore_tools
|
||||||
|
|
||||||
|
# Plan has read-only tools
|
||||||
|
assert "edit_file" not in plan_tools
|
||||||
|
assert "write_file" not in plan_tools
|
||||||
|
|
||||||
|
# Task has all tools plus spawn_agent
|
||||||
|
assert "edit_file" in task_tools
|
||||||
|
assert "write_file" in task_tools
|
||||||
|
assert "spawn_agent" in task_tools
|
||||||
|
|
||||||
|
# Only Task has spawn_agent
|
||||||
|
assert "spawn_agent" not in explore_tools
|
||||||
|
assert "spawn_agent" not in plan_tools
|
||||||
@@ -0,0 +1,90 @@
|
|||||||
|
"""
|
||||||
|
Tests for token counting utilities.
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.shared.tokens import count_tokens, count_message_tokens, estimate_tokens
|
||||||
|
|
||||||
|
|
||||||
|
class TestTokenCounting:
|
||||||
|
"""Tests for token counting functions."""
|
||||||
|
|
||||||
|
def test_estimate_tokens_basic(self):
|
||||||
|
"""Test basic token estimation."""
|
||||||
|
text = "Hello world"
|
||||||
|
tokens = estimate_tokens(text)
|
||||||
|
# ~4 chars per token
|
||||||
|
assert tokens == len(text) // 4
|
||||||
|
|
||||||
|
def test_estimate_tokens_empty(self):
|
||||||
|
"""Test estimation with empty string."""
|
||||||
|
assert estimate_tokens("") == 0
|
||||||
|
|
||||||
|
def test_estimate_tokens_long_text(self):
|
||||||
|
"""Test estimation with longer text."""
|
||||||
|
text = "a" * 400
|
||||||
|
tokens = estimate_tokens(text)
|
||||||
|
assert tokens == 100
|
||||||
|
|
||||||
|
def test_count_tokens_basic(self):
|
||||||
|
"""Test actual token counting."""
|
||||||
|
text = "Hello, how are you today?"
|
||||||
|
tokens = count_tokens(text)
|
||||||
|
# Should return reasonable token count
|
||||||
|
assert tokens > 0
|
||||||
|
assert tokens < len(text) # Should be fewer tokens than characters
|
||||||
|
|
||||||
|
def test_count_tokens_empty(self):
|
||||||
|
"""Test counting empty string."""
|
||||||
|
tokens = count_tokens("")
|
||||||
|
assert tokens == 0
|
||||||
|
|
||||||
|
def test_count_message_tokens_single(self):
|
||||||
|
"""Test counting tokens in single message."""
|
||||||
|
messages = [{"role": "user", "content": "Hello"}]
|
||||||
|
tokens = count_message_tokens(messages)
|
||||||
|
assert tokens > 0
|
||||||
|
|
||||||
|
def test_count_message_tokens_multiple(self):
|
||||||
|
"""Test counting tokens in multiple messages."""
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "Hello, how are you?"},
|
||||||
|
{"role": "assistant", "content": "I'm doing well, thank you!"},
|
||||||
|
]
|
||||||
|
tokens = count_message_tokens(messages)
|
||||||
|
# Should be more than single message
|
||||||
|
single_tokens = count_message_tokens([messages[0]])
|
||||||
|
assert tokens > single_tokens
|
||||||
|
|
||||||
|
def test_count_message_tokens_empty_list(self):
|
||||||
|
"""Test counting empty message list."""
|
||||||
|
tokens = count_message_tokens([])
|
||||||
|
# litellm may return small overhead even for empty list
|
||||||
|
assert tokens < 10
|
||||||
|
|
||||||
|
|
||||||
|
class TestTokenCountingAccuracy:
|
||||||
|
"""Tests for token counting accuracy."""
|
||||||
|
|
||||||
|
def test_code_tokens_reasonable(self):
|
||||||
|
"""Test that code is tokenized reasonably."""
|
||||||
|
code = """
|
||||||
|
def hello_world():
|
||||||
|
print("Hello, World!")
|
||||||
|
return True
|
||||||
|
"""
|
||||||
|
tokens = count_tokens(code)
|
||||||
|
# Code should have reasonable token count
|
||||||
|
assert 10 < tokens < 100
|
||||||
|
|
||||||
|
def test_special_characters(self):
|
||||||
|
"""Test tokenization of special characters."""
|
||||||
|
text = "Hello! @#$%^&*() World?"
|
||||||
|
tokens = count_tokens(text)
|
||||||
|
assert tokens > 0
|
||||||
|
|
||||||
|
def test_unicode_text(self):
|
||||||
|
"""Test tokenization of unicode text."""
|
||||||
|
text = "Hello 世界 🌍"
|
||||||
|
tokens = count_tokens(text)
|
||||||
|
assert tokens > 0
|
||||||
Reference in New Issue
Block a user