ok... ok... I'll add it to git...
This commit is contained in:
@@ -0,0 +1,58 @@
|
||||
# Python
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
*.so
|
||||
.Python
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
|
||||
# Testing
|
||||
.pytest_cache/
|
||||
.coverage
|
||||
htmlcov/
|
||||
.tox/
|
||||
|
||||
# IDEs
|
||||
.vscode/
|
||||
.idea/
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
|
||||
# Logs
|
||||
logs/
|
||||
*.log
|
||||
|
||||
# Environment
|
||||
.env
|
||||
.env.local
|
||||
|
||||
# Git
|
||||
.git/
|
||||
.gitignore
|
||||
|
||||
# Documentation
|
||||
README.md
|
||||
docs/
|
||||
|
||||
# Docker
|
||||
.dockerignore
|
||||
Dockerfile
|
||||
docker-compose.yml
|
||||
@@ -0,0 +1,23 @@
|
||||
# Core Code API Configuration
|
||||
|
||||
# Application settings
|
||||
APP_NAME="Core Code API"
|
||||
APP_VERSION="1.0.0"
|
||||
DEBUG=false
|
||||
|
||||
# Server settings
|
||||
HOST=0.0.0.0
|
||||
PORT=8083
|
||||
|
||||
# CORS settings (default allows all origins for internal use)
|
||||
# CORS_ORIGINS=["http://192.168.86.149:82"]
|
||||
|
||||
# Logging
|
||||
LOG_LEVEL=INFO
|
||||
|
||||
# Web Scraper Module
|
||||
WEB_SCRAPER_REQUEST_TIMEOUT=30
|
||||
WEB_SCRAPER_MAX_REDIRECTS=5
|
||||
WEB_SCRAPER_USER_AGENT="Mozilla/5.0 (compatible; CoreCode/1.0)"
|
||||
WEB_SCRAPER_DEFAULT_MAX_LENGTH=10000
|
||||
WEB_SCRAPER_MAX_LINKS_TO_EXTRACT=50
|
||||
@@ -0,0 +1,47 @@
|
||||
# Use official Python full image with all system tools
|
||||
FROM python:3.12
|
||||
|
||||
# Set working directory
|
||||
WORKDIR /app
|
||||
|
||||
# Set environment variables
|
||||
ENV PYTHONUNBUFFERED=1 \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
PIP_NO_CACHE_DIR=1 \
|
||||
PIP_DISABLE_PIP_VERSION_CHECK=1
|
||||
|
||||
# Install system dependencies
|
||||
RUN apt-get update && \
|
||||
apt-get install -y --no-install-recommends \
|
||||
gcc \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Copy requirements first for better caching
|
||||
COPY requirements.txt .
|
||||
|
||||
# Install Python dependencies
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# Copy application code
|
||||
COPY src/ ./src/
|
||||
|
||||
# Create logs directory
|
||||
RUN mkdir -p logs
|
||||
|
||||
# Create non-root user for security
|
||||
RUN useradd -m -u 1000 appuser && \
|
||||
chown -R appuser:appuser /app
|
||||
|
||||
# Switch to non-root user
|
||||
USER appuser
|
||||
|
||||
# Expose port
|
||||
EXPOSE 8083
|
||||
|
||||
# Health check
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD curl -f http://localhost:8083/health || exit 1
|
||||
|
||||
# Run the application
|
||||
CMD ["uvicorn", "src.main:app", "--host", "0.0.0.0", "--port", "8083"]
|
||||
@@ -0,0 +1,196 @@
|
||||
# Core Code API
|
||||
|
||||
OpenAPI-compatible functions for Open WebUI, providing web scraping and data processing capabilities.
|
||||
|
||||
## Features
|
||||
|
||||
### Web Scraper
|
||||
- Intelligent content extraction using Trafilatura
|
||||
- BeautifulSoup fallback for complex pages
|
||||
- Configurable content length limits
|
||||
- Optional link extraction
|
||||
- Perfect for feeding webpage content to LLMs
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
src/
|
||||
├── config.py # Global application settings
|
||||
├── logging_config.py # Logging configuration
|
||||
├── base_schema.py # Base Pydantic models
|
||||
├── main.py # FastAPI application entry point
|
||||
└── web_scraper/ # Web scraper module
|
||||
├── __init__.py
|
||||
├── config.py # Module-specific settings
|
||||
├── schemas.py # Pydantic request/response models
|
||||
├── service.py # Business logic
|
||||
├── router.py # API routes
|
||||
└── exceptions.py # Custom exceptions
|
||||
```
|
||||
|
||||
## Development
|
||||
|
||||
### Requirements
|
||||
- Python 3.12+
|
||||
- Docker (for containerized deployment)
|
||||
|
||||
### Local Development
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
pip install -r requirements.txt
|
||||
|
||||
# Run locally
|
||||
uvicorn src.main:app --reload --host 0.0.0.0 --port 8083
|
||||
```
|
||||
|
||||
### Docker Build
|
||||
|
||||
```bash
|
||||
# Build image
|
||||
docker build -t core-code:latest .
|
||||
|
||||
# Run container
|
||||
docker run -p 8083:8083 core-code:latest
|
||||
```
|
||||
|
||||
## Deployment
|
||||
|
||||
### Portainer Stack
|
||||
|
||||
1. Navigate to Portainer UI
|
||||
2. Go to **Stacks** → **Add Stack**
|
||||
3. Name: `core-code`
|
||||
4. Upload `stacks/core-code.yml` or paste contents
|
||||
5. Deploy
|
||||
|
||||
### Environment Variables
|
||||
|
||||
See `.env.example` for all available configuration options.
|
||||
|
||||
## API Documentation
|
||||
|
||||
Once deployed, access documentation at:
|
||||
- **Swagger UI**: http://192.168.86.149:8083/docs
|
||||
- **ReDoc**: http://192.168.86.149:8083/redoc
|
||||
- **OpenAPI Spec**: http://192.168.86.149:8083/openapi.json
|
||||
|
||||
## Integration with Open WebUI
|
||||
|
||||
### Method 1: Functions (OpenAPI Import)
|
||||
1. In Open WebUI, navigate to Functions
|
||||
2. Import from OpenAPI spec: `http://192.168.86.149:8083/openapi.json`
|
||||
3. Use functions directly in chat
|
||||
|
||||
### Method 2: Pipelines
|
||||
1. Create a pipeline that calls Core Code API endpoints
|
||||
2. Use as data source for LLM workflows
|
||||
|
||||
### Method 3: Direct API Calls
|
||||
```python
|
||||
import httpx
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(
|
||||
"http://192.168.86.149:8083/web-scraper/scrape",
|
||||
json={
|
||||
"url": "https://example.com",
|
||||
"extract_main_content": True
|
||||
}
|
||||
)
|
||||
data = response.json()
|
||||
```
|
||||
|
||||
## API Endpoints
|
||||
|
||||
### Web Scraper
|
||||
|
||||
**POST /web-scraper/scrape**
|
||||
|
||||
Scrape and extract content from a website.
|
||||
|
||||
Request:
|
||||
```json
|
||||
{
|
||||
"url": "https://example.com/article",
|
||||
"extract_main_content": true,
|
||||
"include_links": false,
|
||||
"max_length": 10000
|
||||
}
|
||||
```
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"url": "https://example.com/article",
|
||||
"title": "Article Title",
|
||||
"content": "Extracted article content...",
|
||||
"extracted_at": "2025-11-12T19:30:00Z",
|
||||
"content_length": 5432,
|
||||
"links": null
|
||||
}
|
||||
```
|
||||
|
||||
## Logging
|
||||
|
||||
Logs are written to:
|
||||
- **Console**: stdout (captured by Docker)
|
||||
- **File**: `/app/logs/app.log` (persisted via volume mount)
|
||||
|
||||
Log format:
|
||||
```
|
||||
2025-11-12 19:30:00 | INFO | src.web_scraper.service:scrape_url:45 | Starting scrape for URL: https://example.com
|
||||
```
|
||||
|
||||
## Health Checks
|
||||
|
||||
- **Endpoint**: `GET /health`
|
||||
- **Docker**: Automatic health checks configured
|
||||
- **Response**: `{"status": "healthy"}`
|
||||
|
||||
## Security
|
||||
|
||||
- Runs as non-root user (uid 1000)
|
||||
- No authentication required (internal network only)
|
||||
- CORS configured for same-network access
|
||||
- Rate limiting: Not implemented (internal use only)
|
||||
|
||||
## Future Modules
|
||||
|
||||
The architecture supports adding new modules:
|
||||
- Data transformation functions
|
||||
- API integrations
|
||||
- File processing
|
||||
- Database queries
|
||||
|
||||
Each module follows the same structure:
|
||||
```
|
||||
src/
|
||||
└── module_name/
|
||||
├── config.py
|
||||
├── schemas.py
|
||||
├── service.py
|
||||
├── router.py
|
||||
└── exceptions.py
|
||||
```
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Container won't start
|
||||
```bash
|
||||
docker logs core-code
|
||||
```
|
||||
|
||||
### API not responding
|
||||
```bash
|
||||
curl http://192.168.86.149:8083/health
|
||||
```
|
||||
|
||||
### Check OpenAPI spec
|
||||
```bash
|
||||
curl http://192.168.86.149:8083/openapi.json | jq
|
||||
```
|
||||
|
||||
## License
|
||||
|
||||
Internal use only.
|
||||
@@ -0,0 +1,222 @@
|
||||
# Core-API Refactoring Plan
|
||||
|
||||
**Date:** 2025-11-14
|
||||
**Goal:** Restructure Core-API into controller-based architecture and add Infrastructure Management API
|
||||
|
||||
## Current Structure
|
||||
|
||||
```
|
||||
src/
|
||||
├── api/
|
||||
│ └── v1/
|
||||
│ ├── chat.py # AI chat completions
|
||||
│ ├── models.py # Model listing
|
||||
│ ├── conversations.py # Conversation memory
|
||||
│ └── schemas.py # Pydantic schemas
|
||||
├── web_scraper/
|
||||
│ ├── router.py # Webscraper endpoints
|
||||
│ ├── service.py
|
||||
│ └── schemas.py
|
||||
├── models/
|
||||
│ ├── ollama_client.py # Ollama HTTP client
|
||||
│ └── embeddings.py
|
||||
├── memory/ # Memory tier system
|
||||
├── config.py # Global settings
|
||||
└── main.py # FastAPI app
|
||||
|
||||
```
|
||||
|
||||
## Target Structure
|
||||
|
||||
```
|
||||
src/
|
||||
├── controllers/ # NEW: Controller-based routing
|
||||
│ ├── __init__.py
|
||||
│ ├── base.py # Base controller class
|
||||
│ ├── ai_controller.py # AI Orchestrator (chat, models, conversations)
|
||||
│ ├── tools_controller.py # Utility tools (webscraper, etc.)
|
||||
│ ├── health_controller.py # Health & monitoring
|
||||
│ └── infrastructure_controller.py # Infrastructure automation
|
||||
├── clients/ # NEW: External API clients
|
||||
│ ├── __init__.py
|
||||
│ ├── portainer_client.py # Portainer API
|
||||
│ ├── npm_client.py # Nginx Proxy Manager API
|
||||
│ └── kuma_client.py # Uptime Kuma Socket.IO API
|
||||
├── api/v1/ # Keep existing for backward compat
|
||||
├── web_scraper/ # Keep as-is for now
|
||||
├── models/ # Keep as-is
|
||||
├── memory/ # Keep as-is
|
||||
├── config.py # Enhanced with infrastructure settings
|
||||
└── main.py # Updated routing
|
||||
|
||||
```
|
||||
|
||||
## Implementation Phases
|
||||
|
||||
### Phase 1: Infrastructure Setup 🔄 IN PROGRESS
|
||||
- [x] Research API authentication methods
|
||||
- [x] Add infrastructure settings to config.py
|
||||
- [ ] Create credentials.py for sensitive data (gitignored)
|
||||
- [ ] Create credentials.example.py as template
|
||||
- [ ] Update .gitignore to exclude credentials.py
|
||||
- [ ] Update config.py to import from credentials module
|
||||
- [x] Create /controllers directory structure
|
||||
- [x] Create /clients directory structure
|
||||
- [x] Create base controller class
|
||||
|
||||
### Phase 2: API Clients ✅ COMPLETE (Portainer & NPM)
|
||||
- [x] Implement Portainer API client (access token auth)
|
||||
- [x] Implement NPM API client (JWT with refresh)
|
||||
- [x] Add token storage/refresh mechanisms
|
||||
- [ ] Implement Uptime Kuma Socket.IO client (DEFERRED - WebSocket complexity)
|
||||
|
||||
### Phase 3: Infrastructure Controller 🔄 IN PROGRESS
|
||||
- [x] GET /infrastructure/health - Check connectivity
|
||||
- [x] GET /infrastructure/services - List all services
|
||||
- [x] GET /infrastructure/services/{name} - Get service details
|
||||
- [x] GET /infrastructure/ports - List allocated ports (skeleton)
|
||||
- [x] GET /infrastructure/domains - List configured domains
|
||||
- [ ] POST /infrastructure/services - Deploy new service
|
||||
- [ ] PUT /infrastructure/services/{name} - Update service
|
||||
- [ ] DELETE /infrastructure/services/{name} - Remove service
|
||||
- [ ] POST /infrastructure/monitoring/add - Auto-add Kuma monitor
|
||||
- [ ] POST /infrastructure/proxy/add - Auto-add NPM proxy host
|
||||
|
||||
### Phase 4: Refactor Existing Controllers 📋 PENDING
|
||||
- [ ] Move AI endpoints to ai_controller.py
|
||||
- [ ] Move webscraper to tools_controller.py
|
||||
- [ ] Move health check to health_controller.py
|
||||
- [ ] Update main.py imports and routing
|
||||
|
||||
### Phase 5: Testing & Documentation 📋 PENDING
|
||||
- [ ] Test all refactored endpoints
|
||||
- [ ] Update API documentation
|
||||
- [ ] Create CLI wrapper scripts
|
||||
- [ ] Remove old shell scripts
|
||||
|
||||
---
|
||||
|
||||
## Progress Notes (2025-11-14)
|
||||
|
||||
### Session 1: Foundation & Read Endpoints
|
||||
**Completed:**
|
||||
- Created controller and client architecture
|
||||
- Implemented Portainer client with full CRUD operations for stacks
|
||||
- Implemented NPM client with JWT refresh and proxy/certificate management
|
||||
- Built infrastructure controller with 5 read/list endpoints
|
||||
- Added infrastructure settings to config.py
|
||||
|
||||
**Files Created:**
|
||||
- `src/controllers/__init__.py`
|
||||
- `src/controllers/base.py`
|
||||
- `src/controllers/infrastructure_controller.py`
|
||||
- `src/clients/__init__.py`
|
||||
- `src/clients/portainer_client.py`
|
||||
- `src/clients/npm_client.py`
|
||||
- `REFACTORING_PLAN.md` (this file)
|
||||
|
||||
**Next Steps:**
|
||||
1. Create credentials.py for secure credential management
|
||||
2. Update config.py to import from credentials module
|
||||
3. Add credentials.py to .gitignore
|
||||
4. Update main.py to include infrastructure routes
|
||||
5. Test endpoints with live infrastructure
|
||||
6. Implement write/deploy operations
|
||||
7. Refactor existing AI/tools/health endpoints
|
||||
8. Create CLI wrappers
|
||||
|
||||
## API Authentication Strategy
|
||||
|
||||
### Portainer
|
||||
- **Method:** Access Token (X-API-Key header)
|
||||
- **Setup:** Manual creation in UI, store in config/env
|
||||
- **Duration:** Long-lived
|
||||
- **Storage:** Environment variable `PORTAINER_API_KEY`
|
||||
|
||||
### Nginx Proxy Manager
|
||||
- **Method:** JWT Bearer Token
|
||||
- **Setup:** Login via `/api/tokens` with credentials
|
||||
- **Duration:** ~24 hours
|
||||
- **Strategy:** Auto-refresh with stored credentials
|
||||
- **Storage:** `NPM_EMAIL` and `NPM_PASSWORD` in env
|
||||
|
||||
### Uptime Kuma
|
||||
- **Method:** Socket.IO WebSocket
|
||||
- **Setup:** Login via Socket.IO `login` event
|
||||
- **Duration:** Session-based
|
||||
- **Strategy:** Maintain persistent connection or re-auth per request
|
||||
- **Storage:** `KUMA_USERNAME` and `KUMA_PASSWORD` in env
|
||||
|
||||
## Configuration Changes
|
||||
|
||||
### Credentials Management Strategy
|
||||
|
||||
**Use `credentials.py` for sensitive data** (added to `.gitignore`):
|
||||
- Keeps secrets out of version control
|
||||
- Easy terminal-based management with editor
|
||||
- Python format for type safety and autocomplete
|
||||
- Separate from config for security isolation
|
||||
|
||||
**Implementation:**
|
||||
1. Create `src/credentials.py` with credentials (gitignored)
|
||||
2. Create `src/credentials.example.py` as template (committed)
|
||||
3. Update `config.py` to import from credentials module
|
||||
4. Add `credentials.py` to `.gitignore`
|
||||
|
||||
**Example `src/credentials.py`:**
|
||||
```python
|
||||
"""
|
||||
Infrastructure credentials (GITIGNORED)
|
||||
Copy from credentials.example.py and fill in real values
|
||||
"""
|
||||
|
||||
# Portainer
|
||||
PORTAINER_URL = "http://localhost:8001"
|
||||
PORTAINER_API_KEY = "ptr_your_actual_token_here"
|
||||
|
||||
# Nginx Proxy Manager
|
||||
NPM_URL = "http://localhost:81"
|
||||
NPM_EMAIL = "jpmschweitzer@gmail.com"
|
||||
NPM_PASSWORD = "your_actual_password"
|
||||
|
||||
# Uptime Kuma
|
||||
KUMA_URL = "http://localhost:3001"
|
||||
KUMA_USERNAME = "admin"
|
||||
KUMA_PASSWORD = "your_actual_password"
|
||||
```
|
||||
|
||||
**Updated `config.py` to use credentials:**
|
||||
```python
|
||||
from src.credentials import (
|
||||
PORTAINER_URL, PORTAINER_API_KEY,
|
||||
NPM_URL, NPM_EMAIL, NPM_PASSWORD,
|
||||
KUMA_URL, KUMA_USERNAME, KUMA_PASSWORD
|
||||
)
|
||||
|
||||
class Settings(BaseSettings):
|
||||
# Infrastructure Management (from credentials.py)
|
||||
portainer_url: str = PORTAINER_URL
|
||||
portainer_api_key: str = PORTAINER_API_KEY
|
||||
|
||||
npm_url: str = NPM_URL
|
||||
npm_email: str = NPM_EMAIL
|
||||
npm_password: str = NPM_PASSWORD
|
||||
|
||||
kuma_url: str = KUMA_URL
|
||||
kuma_username: str = KUMA_USERNAME
|
||||
kuma_password: str = KUMA_PASSWORD
|
||||
```
|
||||
|
||||
## Benefits
|
||||
|
||||
1. **Cleaner Code:** Separation of concerns, easier to maintain
|
||||
2. **Automation:** Programmatic service deployment and configuration
|
||||
3. **Elimination of Shell Scripts:** Replace ad-hoc scripts with proper API
|
||||
4. **Service Discovery:** Auto-detect running services and configurations
|
||||
5. **Self-Managing Homelab:** Foundation for autonomous infrastructure
|
||||
|
||||
## Migration Notes
|
||||
|
||||
- Existing `/v1/` endpoints remain unchanged for backward compatibility
|
||||
- Web scraper endpoints stay at `/web-scraper/` initially
|
||||
- Old shell scripts in `/stacks/` will be replaced with CLI wrappers
|
||||
@@ -0,0 +1,22 @@
|
||||
# FastAPI and ASGI server
|
||||
fastapi==0.115.0
|
||||
uvicorn[standard]==0.32.0
|
||||
pydantic==2.10.4
|
||||
pydantic-settings==2.7.0
|
||||
|
||||
# HTTP client
|
||||
httpx==0.28.1
|
||||
|
||||
# Web scraping
|
||||
beautifulsoup4==4.12.3
|
||||
trafilatura==1.12.2
|
||||
lxml==5.3.0
|
||||
|
||||
# Utilities
|
||||
python-multipart==0.0.12
|
||||
python-dotenv==1.0.1
|
||||
python-json-logger==2.0.7
|
||||
|
||||
# Memory & Embeddings
|
||||
qdrant-client==1.11.3
|
||||
sentence-transformers==3.3.1
|
||||
@@ -0,0 +1,6 @@
|
||||
"""
|
||||
Core Code API - OpenAPI-compatible functions for Open WebUI
|
||||
"""
|
||||
|
||||
__version__ = "1.0.0"
|
||||
__author__ = "Core Code Team"
|
||||
@@ -0,0 +1,291 @@
|
||||
"""
|
||||
OpenAI-compatible /v1/chat/completions endpoint
|
||||
|
||||
Phase 2: Integrated with memory system for conversation persistence.
|
||||
"""
|
||||
|
||||
import time
|
||||
import logging
|
||||
import uuid
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from typing import AsyncIterator
|
||||
|
||||
from .schemas import (
|
||||
ChatCompletionRequest,
|
||||
ChatCompletionResponse,
|
||||
ChatCompletionChoice,
|
||||
ChatMessageResponse,
|
||||
UsageInfo,
|
||||
ChatCompletionStreamResponse,
|
||||
ChatCompletionStreamChoice,
|
||||
DeltaMessage,
|
||||
)
|
||||
from src.models.ollama_client import get_ollama_client
|
||||
from src.memory import get_memory_manager, MessageRole as MemoryMessageRole, TokenUsage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def build_prompt_from_messages(messages: list) -> str:
|
||||
"""
|
||||
Convert message list to a prompt string.
|
||||
|
||||
In Phase 1, we do simple concatenation.
|
||||
Phase 2 will add proper memory management.
|
||||
"""
|
||||
prompt_parts = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.role.value if hasattr(msg.role, 'value') else msg.role
|
||||
content = msg.content
|
||||
|
||||
if role == "system":
|
||||
prompt_parts.append(f"System: {content}")
|
||||
elif role == "user":
|
||||
prompt_parts.append(f"User: {content}")
|
||||
elif role == "assistant":
|
||||
prompt_parts.append(f"Assistant: {content}")
|
||||
|
||||
prompt_parts.append("Assistant:")
|
||||
return "\n\n".join(prompt_parts)
|
||||
|
||||
|
||||
async def stream_chat_completion(
|
||||
request_id: str,
|
||||
model: str,
|
||||
prompt: str,
|
||||
temperature: float,
|
||||
max_tokens: int | None
|
||||
) -> AsyncIterator[str]:
|
||||
"""
|
||||
Stream chat completion in OpenAI SSE format.
|
||||
|
||||
Yields:
|
||||
Server-Sent Events formatted strings
|
||||
"""
|
||||
created = int(time.time())
|
||||
ollama_client = get_ollama_client()
|
||||
|
||||
# First chunk with role
|
||||
first_chunk = ChatCompletionStreamResponse(
|
||||
id=request_id,
|
||||
created=created,
|
||||
model=model,
|
||||
choices=[
|
||||
ChatCompletionStreamChoice(
|
||||
index=0,
|
||||
delta=DeltaMessage(role="assistant"),
|
||||
finish_reason=None
|
||||
)
|
||||
]
|
||||
)
|
||||
yield f"data: {first_chunk.model_dump_json()}\n\n"
|
||||
|
||||
# Stream tokens
|
||||
try:
|
||||
async for token in ollama_client.generate_streaming(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens
|
||||
):
|
||||
chunk = ChatCompletionStreamResponse(
|
||||
id=request_id,
|
||||
created=created,
|
||||
model=model,
|
||||
choices=[
|
||||
ChatCompletionStreamChoice(
|
||||
index=0,
|
||||
delta=DeltaMessage(content=token),
|
||||
finish_reason=None
|
||||
)
|
||||
]
|
||||
)
|
||||
yield f"data: {chunk.model_dump_json()}\n\n"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Streaming error: {e}")
|
||||
# Send error in OpenAI format
|
||||
import json
|
||||
error_chunk = {
|
||||
"error": {
|
||||
"message": str(e),
|
||||
"type": "server_error"
|
||||
}
|
||||
}
|
||||
yield f"data: {json.dumps(error_chunk)}\n\n"
|
||||
return
|
||||
|
||||
# Final chunk
|
||||
final_chunk = ChatCompletionStreamResponse(
|
||||
id=request_id,
|
||||
created=created,
|
||||
model=model,
|
||||
choices=[
|
||||
ChatCompletionStreamChoice(
|
||||
index=0,
|
||||
delta=DeltaMessage(),
|
||||
finish_reason="stop"
|
||||
)
|
||||
]
|
||||
)
|
||||
yield f"data: {final_chunk.model_dump_json()}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
|
||||
async def store_conversation_turn(
|
||||
conversation_id: str,
|
||||
role: str,
|
||||
content: str,
|
||||
tokens: dict = None
|
||||
):
|
||||
"""
|
||||
Store a conversation turn in memory
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
role: Message role (user, assistant, system)
|
||||
content: Message content
|
||||
tokens: Optional token usage dict
|
||||
"""
|
||||
try:
|
||||
memory_manager = get_memory_manager()
|
||||
|
||||
# Convert role string to MemoryMessageRole
|
||||
if role == "user":
|
||||
memory_role = MemoryMessageRole.USER
|
||||
elif role == "assistant":
|
||||
memory_role = MemoryMessageRole.ASSISTANT
|
||||
elif role == "system":
|
||||
memory_role = MemoryMessageRole.SYSTEM
|
||||
else:
|
||||
memory_role = MemoryMessageRole.USER # Default fallback
|
||||
|
||||
# Create TokenUsage if provided
|
||||
token_usage = None
|
||||
if tokens:
|
||||
token_usage = TokenUsage(
|
||||
prompt=tokens.get("prompt", 0),
|
||||
completion=tokens.get("completion", 0),
|
||||
total=tokens.get("total", 0)
|
||||
)
|
||||
|
||||
# Store in memory
|
||||
await memory_manager.add_turn(
|
||||
conversation_id=conversation_id,
|
||||
role=memory_role,
|
||||
content=content,
|
||||
tokens=token_usage
|
||||
)
|
||||
|
||||
logger.debug(f"Stored {role} turn in memory for conversation {conversation_id}")
|
||||
|
||||
except Exception as e:
|
||||
# Log error but don't fail the request
|
||||
logger.error(f"Failed to store turn in memory: {e}")
|
||||
|
||||
|
||||
@router.post("/v1/chat/completions")
|
||||
async def chat_completions(request: ChatCompletionRequest):
|
||||
"""
|
||||
OpenAI-compatible chat completions endpoint.
|
||||
Supports both streaming and non-streaming.
|
||||
|
||||
Phase 2: Automatically stores conversations in memory system.
|
||||
"""
|
||||
request_id = f"chatcmpl-{int(time.time() * 1000)}"
|
||||
|
||||
# Generate or use provided conversation_id
|
||||
conversation_id = request.conversation_id or f"conv_{uuid.uuid4().hex[:16]}"
|
||||
|
||||
logger.info(
|
||||
f"Chat request: id={request_id}, model={request.model}, "
|
||||
f"messages={len(request.messages)}, stream={request.stream}, "
|
||||
f"conversation_id={conversation_id}, store_in_memory={request.store_in_memory}"
|
||||
)
|
||||
|
||||
# Store user messages in memory (if enabled)
|
||||
if request.store_in_memory:
|
||||
for msg in request.messages:
|
||||
role = msg.role.value if hasattr(msg.role, 'value') else msg.role
|
||||
if role == "user": # Store latest user message
|
||||
await store_conversation_turn(
|
||||
conversation_id=conversation_id,
|
||||
role=role,
|
||||
content=msg.content
|
||||
)
|
||||
|
||||
# Build prompt from messages
|
||||
prompt = build_prompt_from_messages(request.messages)
|
||||
|
||||
# Streaming response
|
||||
if request.stream:
|
||||
return StreamingResponse(
|
||||
stream_chat_completion(
|
||||
request_id=request_id,
|
||||
model=request.model,
|
||||
prompt=prompt,
|
||||
temperature=request.temperature,
|
||||
max_tokens=request.max_tokens
|
||||
),
|
||||
media_type="text/event-stream"
|
||||
)
|
||||
|
||||
# Non-streaming response
|
||||
try:
|
||||
ollama_client = get_ollama_client()
|
||||
result = await ollama_client.generate_non_streaming(
|
||||
model=request.model,
|
||||
prompt=prompt,
|
||||
temperature=request.temperature,
|
||||
max_tokens=request.max_tokens
|
||||
)
|
||||
|
||||
assistant_content = result["response"]
|
||||
|
||||
# Store assistant response in memory (if enabled)
|
||||
if request.store_in_memory:
|
||||
await store_conversation_turn(
|
||||
conversation_id=conversation_id,
|
||||
role="assistant",
|
||||
content=assistant_content,
|
||||
tokens=result["tokens"]
|
||||
)
|
||||
|
||||
response = ChatCompletionResponse(
|
||||
id=request_id,
|
||||
created=int(time.time()),
|
||||
model=request.model,
|
||||
choices=[
|
||||
ChatCompletionChoice(
|
||||
index=0,
|
||||
message=ChatMessageResponse(
|
||||
role="assistant",
|
||||
content=assistant_content
|
||||
),
|
||||
finish_reason="stop"
|
||||
)
|
||||
],
|
||||
usage=UsageInfo(
|
||||
prompt_tokens=result["tokens"]["prompt"],
|
||||
completion_tokens=result["tokens"]["completion"],
|
||||
total_tokens=result["tokens"]["total"]
|
||||
)
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"Chat response: id={request_id}, "
|
||||
f"tokens={result['tokens']['total']}, "
|
||||
f"conversation_id={conversation_id}"
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Chat completion error: {e}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Failed to generate completion: {str(e)}"
|
||||
)
|
||||
@@ -0,0 +1,338 @@
|
||||
"""
|
||||
Conversation History API Endpoints
|
||||
|
||||
Provides endpoints for managing and querying conversation memory:
|
||||
- List conversations
|
||||
- Get conversation history
|
||||
- Search conversations semantically
|
||||
- Delete conversations
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException, Query
|
||||
from typing import List, Optional
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from src.memory import get_memory_manager, MessageRole
|
||||
|
||||
router = APIRouter(prefix="/v1/conversations", tags=["conversations"])
|
||||
|
||||
|
||||
# Request/Response Models
|
||||
class SearchRequest(BaseModel):
|
||||
"""Request model for semantic search"""
|
||||
query: str = Field(..., description="Search query")
|
||||
limit: int = Field(5, ge=1, le=50, description="Maximum number of results")
|
||||
|
||||
|
||||
class ConversationTurnResponse(BaseModel):
|
||||
"""Response model for a conversation turn"""
|
||||
turn_number: int
|
||||
role: str
|
||||
content: str
|
||||
timestamp: str
|
||||
tokens_prompt: Optional[int] = None
|
||||
tokens_completion: Optional[int] = None
|
||||
tokens_total: Optional[int] = None
|
||||
metadata: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
class ConversationHistoryResponse(BaseModel):
|
||||
"""Response model for conversation history"""
|
||||
conversation_id: str
|
||||
turn_count: int
|
||||
total_tokens: int
|
||||
turns: List[ConversationTurnResponse]
|
||||
|
||||
|
||||
class SearchResultResponse(BaseModel):
|
||||
"""Response model for a single search result"""
|
||||
conversation_id: str
|
||||
turn_number: int
|
||||
role: str
|
||||
content: str
|
||||
timestamp: str
|
||||
score: float
|
||||
|
||||
|
||||
class SearchResponse(BaseModel):
|
||||
"""Response model for search results"""
|
||||
query: str
|
||||
results: List[SearchResultResponse]
|
||||
count: int
|
||||
|
||||
|
||||
class ConversationStatsResponse(BaseModel):
|
||||
"""Response model for conversation statistics"""
|
||||
conversation_id: str
|
||||
buffer_turns: int
|
||||
buffer_tokens: int
|
||||
qdrant_turns: int
|
||||
qdrant_tokens: int
|
||||
exists_in_buffer: bool
|
||||
exists_in_qdrant: bool
|
||||
|
||||
|
||||
class DeleteResponse(BaseModel):
|
||||
"""Response model for delete operation"""
|
||||
conversation_id: str
|
||||
deleted: bool
|
||||
message: str
|
||||
|
||||
|
||||
# Endpoints
|
||||
|
||||
@router.get(
|
||||
"/{conversation_id}",
|
||||
response_model=ConversationHistoryResponse,
|
||||
summary="Get conversation history",
|
||||
description="Retrieve complete conversation history including all turns"
|
||||
)
|
||||
async def get_conversation(
|
||||
conversation_id: str,
|
||||
include_buffer: bool = Query(
|
||||
True,
|
||||
description="Include recent turns from buffer that haven't been consolidated yet"
|
||||
)
|
||||
):
|
||||
"""
|
||||
Get complete conversation history
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
include_buffer: Include recent buffer turns not yet consolidated
|
||||
|
||||
Returns:
|
||||
Complete conversation history with all turns
|
||||
"""
|
||||
manager = get_memory_manager()
|
||||
|
||||
# Get full history
|
||||
turns = await manager.get_full_history(conversation_id, include_buffer=include_buffer)
|
||||
|
||||
if not turns:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Conversation {conversation_id} not found"
|
||||
)
|
||||
|
||||
# Convert to response format
|
||||
turn_responses = []
|
||||
total_tokens = 0
|
||||
|
||||
for turn in turns:
|
||||
turn_response = ConversationTurnResponse(
|
||||
turn_number=turn.turn_number,
|
||||
role=turn.role.value if isinstance(turn.role, MessageRole) else turn.role,
|
||||
content=turn.content,
|
||||
timestamp=turn.timestamp.isoformat(),
|
||||
metadata=turn.metadata
|
||||
)
|
||||
|
||||
if turn.tokens:
|
||||
turn_response.tokens_prompt = turn.tokens.prompt
|
||||
turn_response.tokens_completion = turn.tokens.completion
|
||||
turn_response.tokens_total = turn.tokens.total
|
||||
total_tokens += turn.tokens.total
|
||||
|
||||
turn_responses.append(turn_response)
|
||||
|
||||
return ConversationHistoryResponse(
|
||||
conversation_id=conversation_id,
|
||||
turn_count=len(turns),
|
||||
total_tokens=total_tokens,
|
||||
turns=turn_responses
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/{conversation_id}/stats",
|
||||
response_model=ConversationStatsResponse,
|
||||
summary="Get conversation statistics",
|
||||
description="Get detailed statistics about a conversation across all storage tiers"
|
||||
)
|
||||
async def get_conversation_stats(conversation_id: str):
|
||||
"""
|
||||
Get conversation statistics
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
Statistics including turn counts and token usage across tiers
|
||||
"""
|
||||
manager = get_memory_manager()
|
||||
stats = await manager.get_conversation_stats(conversation_id)
|
||||
|
||||
return ConversationStatsResponse(**stats)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{conversation_id}/search",
|
||||
response_model=SearchResponse,
|
||||
summary="Search conversation semantically",
|
||||
description="Search for relevant turns within a conversation using semantic similarity"
|
||||
)
|
||||
async def search_conversation(
|
||||
conversation_id: str,
|
||||
search_request: SearchRequest
|
||||
):
|
||||
"""
|
||||
Semantic search within a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
search_request: Search query and parameters
|
||||
|
||||
Returns:
|
||||
Relevant conversation turns ranked by semantic similarity
|
||||
"""
|
||||
manager = get_memory_manager()
|
||||
|
||||
# Perform semantic search
|
||||
results = await manager.search_conversations(
|
||||
query=search_request.query,
|
||||
conversation_id=conversation_id,
|
||||
limit=search_request.limit
|
||||
)
|
||||
|
||||
# Convert to response format
|
||||
search_results = [
|
||||
SearchResultResponse(
|
||||
conversation_id=result["conversation_id"],
|
||||
turn_number=result["turn_number"],
|
||||
role=result["role"],
|
||||
content=result["content"],
|
||||
timestamp=result["timestamp"],
|
||||
score=result["score"]
|
||||
)
|
||||
for result in results
|
||||
]
|
||||
|
||||
return SearchResponse(
|
||||
query=search_request.query,
|
||||
results=search_results,
|
||||
count=len(search_results)
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/search",
|
||||
response_model=SearchResponse,
|
||||
summary="Search all conversations",
|
||||
description="Search across all conversations using semantic similarity"
|
||||
)
|
||||
async def search_all_conversations(search_request: SearchRequest):
|
||||
"""
|
||||
Semantic search across all conversations
|
||||
|
||||
Args:
|
||||
search_request: Search query and parameters
|
||||
|
||||
Returns:
|
||||
Relevant turns from any conversation ranked by semantic similarity
|
||||
"""
|
||||
manager = get_memory_manager()
|
||||
|
||||
# Perform semantic search across all conversations
|
||||
results = await manager.search_conversations(
|
||||
query=search_request.query,
|
||||
conversation_id=None, # Search all conversations
|
||||
limit=search_request.limit
|
||||
)
|
||||
|
||||
# Convert to response format
|
||||
search_results = [
|
||||
SearchResultResponse(
|
||||
conversation_id=result["conversation_id"],
|
||||
turn_number=result["turn_number"],
|
||||
role=result["role"],
|
||||
content=result["content"],
|
||||
timestamp=result["timestamp"],
|
||||
score=result["score"]
|
||||
)
|
||||
for result in results
|
||||
]
|
||||
|
||||
return SearchResponse(
|
||||
query=search_request.query,
|
||||
results=search_results,
|
||||
count=len(search_results)
|
||||
)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/{conversation_id}",
|
||||
response_model=DeleteResponse,
|
||||
summary="Delete conversation",
|
||||
description="Delete a conversation from all storage tiers"
|
||||
)
|
||||
async def delete_conversation(
|
||||
conversation_id: str,
|
||||
clear_buffer: bool = Query(True, description="Clear from buffer (Tier 1)"),
|
||||
clear_qdrant: bool = Query(True, description="Clear from Qdrant (Tier 2/3)")
|
||||
):
|
||||
"""
|
||||
Delete a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
clear_buffer: Clear from Tier 1 buffer
|
||||
clear_qdrant: Clear from Tier 2/3 Qdrant
|
||||
|
||||
Returns:
|
||||
Deletion confirmation
|
||||
"""
|
||||
manager = get_memory_manager()
|
||||
|
||||
try:
|
||||
await manager.clear_conversation(
|
||||
conversation_id,
|
||||
clear_buffer=clear_buffer,
|
||||
clear_qdrant=clear_qdrant
|
||||
)
|
||||
|
||||
return DeleteResponse(
|
||||
conversation_id=conversation_id,
|
||||
deleted=True,
|
||||
message=f"Conversation {conversation_id} deleted successfully"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Error deleting conversation: {str(e)}"
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{conversation_id}/consolidate",
|
||||
summary="Consolidate conversation",
|
||||
description="Manually trigger consolidation from buffer to persistent storage"
|
||||
)
|
||||
async def consolidate_conversation(conversation_id: str):
|
||||
"""
|
||||
Manually consolidate a conversation
|
||||
|
||||
Moves all buffer turns to Qdrant persistent storage.
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
Number of turns consolidated
|
||||
"""
|
||||
manager = get_memory_manager()
|
||||
|
||||
try:
|
||||
count = await manager.consolidate(conversation_id)
|
||||
|
||||
return {
|
||||
"conversation_id": conversation_id,
|
||||
"consolidated_turns": count,
|
||||
"message": f"Successfully consolidated {count} turns to persistent storage"
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Error consolidating conversation: {str(e)}"
|
||||
)
|
||||
@@ -0,0 +1,34 @@
|
||||
"""
|
||||
OpenAI-compatible /v1/models endpoint
|
||||
"""
|
||||
|
||||
from fastapi import APIRouter
|
||||
from .schemas import ModelsListResponse, ModelInfo
|
||||
from src.config import get_settings
|
||||
|
||||
router = APIRouter()
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
@router.get("/v1/models")
|
||||
async def list_models():
|
||||
"""List available models in OpenAI format."""
|
||||
|
||||
models = []
|
||||
|
||||
# Add OpenAI-style aliases
|
||||
for alias in settings.model_aliases.keys():
|
||||
models.append(ModelInfo(id=alias, owned_by="tatlock"))
|
||||
|
||||
# Add actual local models
|
||||
for model_list in [
|
||||
settings.get_lightweight_models(),
|
||||
settings.get_heavy_models(),
|
||||
settings.get_code_models()
|
||||
]:
|
||||
for model in model_list:
|
||||
# Avoid duplicates
|
||||
if model not in [m.id for m in models]:
|
||||
models.append(ModelInfo(id=model, owned_by="tatlock"))
|
||||
|
||||
return ModelsListResponse(data=models)
|
||||
@@ -0,0 +1,142 @@
|
||||
"""
|
||||
OpenAI-compatible API schemas for /v1/* endpoints
|
||||
"""
|
||||
from pydantic import BaseModel, Field
|
||||
from typing import List, Optional, Literal
|
||||
from enum import Enum
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Request Schemas
|
||||
# ============================================================================
|
||||
|
||||
class MessageRole(str, Enum):
|
||||
"""Valid message roles."""
|
||||
SYSTEM = "system"
|
||||
USER = "user"
|
||||
ASSISTANT = "assistant"
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
"""A single message in the conversation."""
|
||||
role: MessageRole
|
||||
content: str
|
||||
name: Optional[str] = None
|
||||
|
||||
|
||||
class ChatCompletionRequest(BaseModel):
|
||||
"""OpenAI-compatible chat completion request."""
|
||||
|
||||
model: str = Field(..., description="Model to use")
|
||||
messages: List[ChatMessage] = Field(..., min_length=1)
|
||||
stream: bool = Field(default=False, description="Enable streaming")
|
||||
|
||||
# Memory system (Phase 2)
|
||||
conversation_id: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Conversation ID for memory tracking (auto-generated if not provided)"
|
||||
)
|
||||
store_in_memory: bool = Field(
|
||||
default=True,
|
||||
description="Store conversation turns in memory system"
|
||||
)
|
||||
|
||||
# Optional parameters
|
||||
temperature: Optional[float] = Field(default=0.7, ge=0, le=2)
|
||||
top_p: Optional[float] = Field(default=1.0, ge=0, le=1)
|
||||
max_tokens: Optional[int] = Field(default=None, ge=1)
|
||||
frequency_penalty: Optional[float] = Field(default=0.0, ge=-2, le=2)
|
||||
presence_penalty: Optional[float] = Field(default=0.0, ge=-2, le=2)
|
||||
stop: Optional[List[str]] = None
|
||||
|
||||
class Config:
|
||||
json_schema_extra = {
|
||||
"example": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello!"}
|
||||
],
|
||||
"stream": False,
|
||||
"temperature": 0.7
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Response Schemas
|
||||
# ============================================================================
|
||||
|
||||
class ChatMessageResponse(BaseModel):
|
||||
"""Response message."""
|
||||
role: str = "assistant"
|
||||
content: str
|
||||
|
||||
|
||||
class ChatCompletionChoice(BaseModel):
|
||||
"""A single completion choice."""
|
||||
index: int = 0
|
||||
message: ChatMessageResponse
|
||||
finish_reason: str = "stop"
|
||||
|
||||
|
||||
class UsageInfo(BaseModel):
|
||||
"""Token usage information."""
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
total_tokens: int = 0
|
||||
|
||||
|
||||
class ChatCompletionResponse(BaseModel):
|
||||
"""OpenAI-compatible chat completion response (non-streaming)."""
|
||||
|
||||
id: str
|
||||
object: str = "chat.completion"
|
||||
created: int
|
||||
model: str
|
||||
choices: List[ChatCompletionChoice]
|
||||
usage: UsageInfo
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Streaming Response Schemas
|
||||
# ============================================================================
|
||||
|
||||
class DeltaMessage(BaseModel):
|
||||
"""Delta message for streaming."""
|
||||
role: Optional[str] = None
|
||||
content: Optional[str] = None
|
||||
|
||||
|
||||
class ChatCompletionStreamChoice(BaseModel):
|
||||
"""Streaming choice."""
|
||||
index: int = 0
|
||||
delta: DeltaMessage
|
||||
finish_reason: Optional[str] = None
|
||||
|
||||
|
||||
class ChatCompletionStreamResponse(BaseModel):
|
||||
"""OpenAI-compatible streaming chunk."""
|
||||
|
||||
id: str
|
||||
object: str = "chat.completion.chunk"
|
||||
created: int
|
||||
model: str
|
||||
choices: List[ChatCompletionStreamChoice]
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Models Endpoint
|
||||
# ============================================================================
|
||||
|
||||
class ModelInfo(BaseModel):
|
||||
"""Model information."""
|
||||
id: str
|
||||
object: str = "model"
|
||||
created: int = 0
|
||||
owned_by: str = "local"
|
||||
|
||||
|
||||
class ModelsListResponse(BaseModel):
|
||||
"""List of available models."""
|
||||
object: str = "list"
|
||||
data: List[ModelInfo]
|
||||
@@ -0,0 +1,45 @@
|
||||
"""
|
||||
Base Pydantic models for consistent schema behavior
|
||||
"""
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
|
||||
class BaseSchema(BaseModel):
|
||||
"""
|
||||
Base Pydantic model with standardized configuration
|
||||
|
||||
All schemas should inherit from this to ensure consistent behavior:
|
||||
- Consistent datetime serialization
|
||||
- Strict validation by default
|
||||
- JSON schema generation
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
# Strict type validation
|
||||
strict=False,
|
||||
|
||||
# Allow population by field name
|
||||
populate_by_name=True,
|
||||
|
||||
# Use enum values in JSON
|
||||
use_enum_values=True,
|
||||
|
||||
# Validate assignments after initialization
|
||||
validate_assignment=True,
|
||||
|
||||
# Serialize datetime to ISO format
|
||||
json_encoders={
|
||||
datetime: lambda v: v.isoformat() if v else None
|
||||
}
|
||||
)
|
||||
|
||||
def dict_without_none(self) -> dict[str, Any]:
|
||||
"""
|
||||
Return model as dict, excluding None values
|
||||
|
||||
Returns:
|
||||
Dictionary with None values filtered out
|
||||
"""
|
||||
return {k: v for k, v in self.model_dump().items() if v is not None}
|
||||
@@ -0,0 +1,5 @@
|
||||
"""
|
||||
API Clients package for Core-API
|
||||
|
||||
Provides HTTP/WebSocket clients for external infrastructure services.
|
||||
"""
|
||||
@@ -0,0 +1,271 @@
|
||||
"""
|
||||
Nginx Proxy Manager API Client
|
||||
|
||||
Provides interface to NPM REST API for proxy host and SSL certificate management.
|
||||
"""
|
||||
import httpx
|
||||
from typing import Optional, Dict, List, Any
|
||||
from datetime import datetime, timedelta
|
||||
from src.logging_config import get_logger
|
||||
from src.config import get_settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class NPMClient:
|
||||
"""
|
||||
HTTP client for Nginx Proxy Manager API
|
||||
|
||||
Uses JWT Bearer token authentication with automatic token refresh.
|
||||
Tokens expire after ~24 hours.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
email: Optional[str] = None,
|
||||
password: Optional[str] = None,
|
||||
timeout: int = 30
|
||||
):
|
||||
"""
|
||||
Initialize NPM client
|
||||
|
||||
Args:
|
||||
base_url: NPM base URL (default from settings)
|
||||
email: NPM admin email (default from settings)
|
||||
password: NPM admin password (default from settings)
|
||||
timeout: Request timeout in seconds
|
||||
"""
|
||||
self.base_url = (base_url or settings.npm_url).rstrip("/")
|
||||
self.email = email or settings.npm_email
|
||||
self.password = password or settings.npm_password
|
||||
self.timeout = timeout
|
||||
|
||||
self._token: Optional[str] = None
|
||||
self._token_expires: Optional[datetime] = None
|
||||
|
||||
if not self.email or not self.password:
|
||||
logger.warning("NPM credentials not configured")
|
||||
|
||||
async def _ensure_token(self):
|
||||
"""Ensure we have a valid token, refresh if needed"""
|
||||
if self._token and self._token_expires:
|
||||
# If token expires in less than 1 hour, refresh it
|
||||
if datetime.now() + timedelta(hours=1) < self._token_expires:
|
||||
return
|
||||
|
||||
# Get new token
|
||||
await self._refresh_token()
|
||||
|
||||
async def _refresh_token(self):
|
||||
"""Get a new authentication token"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/tokens",
|
||||
json={
|
||||
"identity": self.email,
|
||||
"secret": self.password
|
||||
}
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
self._token = data.get("token")
|
||||
# Assume 23-hour expiration to be safe
|
||||
self._token_expires = datetime.now() + timedelta(hours=23)
|
||||
|
||||
logger.info("NPM token refreshed successfully")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to refresh NPM token: {e}")
|
||||
raise
|
||||
|
||||
def _get_headers(self) -> Dict[str, str]:
|
||||
"""Get request headers with authentication"""
|
||||
if not self._token:
|
||||
raise RuntimeError("No NPM token available. Call _ensure_token() first.")
|
||||
|
||||
return {
|
||||
"Authorization": f"Bearer {self._token}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""
|
||||
Check if NPM API is accessible
|
||||
|
||||
Returns:
|
||||
True if accessible, False otherwise
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(f"{self.base_url}/api")
|
||||
return response.status_code == 200
|
||||
except Exception as e:
|
||||
logger.error(f"NPM health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def get_proxy_hosts(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all proxy hosts
|
||||
|
||||
Returns:
|
||||
List of proxy host configurations
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/nginx/proxy-hosts",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_proxy_host(self, host_id: int) -> Dict[str, Any]:
|
||||
"""
|
||||
Get details of a specific proxy host
|
||||
|
||||
Args:
|
||||
host_id: Proxy host identifier
|
||||
|
||||
Returns:
|
||||
Proxy host configuration
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/nginx/proxy-hosts/{host_id}",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def create_proxy_host(
|
||||
self,
|
||||
domain_names: List[str],
|
||||
forward_host: str,
|
||||
forward_port: int,
|
||||
forward_scheme: str = "http",
|
||||
certificate_id: int = 0,
|
||||
ssl_forced: bool = False,
|
||||
block_exploits: bool = True,
|
||||
caching_enabled: bool = True,
|
||||
websocket_upgrade: bool = True,
|
||||
http2_support: bool = True,
|
||||
hsts_enabled: bool = True,
|
||||
advanced_config: str = ""
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create a new proxy host
|
||||
|
||||
Args:
|
||||
domain_names: List of domain names for this proxy
|
||||
forward_host: Target host to proxy to
|
||||
forward_port: Target port to proxy to
|
||||
forward_scheme: http or https
|
||||
certificate_id: SSL certificate ID (0 for none)
|
||||
ssl_forced: Force HTTPS redirect
|
||||
block_exploits: Enable exploit blocking
|
||||
caching_enabled: Enable response caching
|
||||
websocket_upgrade: Allow WebSocket upgrades
|
||||
http2_support: Enable HTTP/2
|
||||
hsts_enabled: Enable HSTS headers
|
||||
advanced_config: Custom nginx configuration
|
||||
|
||||
Returns:
|
||||
Created proxy host details
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
payload = {
|
||||
"domain_names": domain_names,
|
||||
"forward_scheme": forward_scheme,
|
||||
"forward_host": forward_host,
|
||||
"forward_port": forward_port,
|
||||
"certificate_id": certificate_id,
|
||||
"ssl_forced": ssl_forced,
|
||||
"block_exploits": block_exploits,
|
||||
"caching_enabled": caching_enabled,
|
||||
"allow_websocket_upgrade": websocket_upgrade,
|
||||
"http2_support": http2_support,
|
||||
"hsts_enabled": hsts_enabled,
|
||||
"hsts_subdomains": False,
|
||||
"advanced_config": advanced_config,
|
||||
"access_list_id": 0,
|
||||
"meta": {}
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/nginx/proxy-hosts",
|
||||
headers=self._get_headers(),
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_certificates(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all SSL certificates
|
||||
|
||||
Returns:
|
||||
List of certificate details
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/nginx/certificates",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def create_certificate(
|
||||
self,
|
||||
domain_names: List[str],
|
||||
provider: str = "letsencrypt"
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Request a new SSL certificate from Let's Encrypt
|
||||
|
||||
Args:
|
||||
domain_names: List of domains for the certificate
|
||||
provider: Certificate provider (default: letsencrypt)
|
||||
|
||||
Returns:
|
||||
Certificate details
|
||||
"""
|
||||
await self._ensure_token()
|
||||
|
||||
payload = {
|
||||
"provider": provider,
|
||||
"domain_names": domain_names,
|
||||
"meta": {
|
||||
"dns_challenge": False
|
||||
}
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/nginx/certificates",
|
||||
headers=self._get_headers(),
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_npm_client: Optional[NPMClient] = None
|
||||
|
||||
|
||||
def get_npm_client() -> NPMClient:
|
||||
"""Get singleton NPM client instance"""
|
||||
global _npm_client
|
||||
if _npm_client is None:
|
||||
_npm_client = NPMClient()
|
||||
return _npm_client
|
||||
@@ -0,0 +1,221 @@
|
||||
"""
|
||||
Portainer API Client
|
||||
|
||||
Provides interface to Portainer REST API for stack and container management.
|
||||
"""
|
||||
import httpx
|
||||
from typing import Optional, Dict, List, Any
|
||||
from src.logging_config import get_logger
|
||||
from src.config import get_settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class PortainerClient:
|
||||
"""
|
||||
HTTP client for Portainer API
|
||||
|
||||
Uses access token authentication (X-API-Key header)
|
||||
for long-lived API access without session management.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
timeout: int = 30
|
||||
):
|
||||
"""
|
||||
Initialize Portainer client
|
||||
|
||||
Args:
|
||||
base_url: Portainer base URL (default from settings)
|
||||
api_key: Portainer API access token (default from settings)
|
||||
timeout: Request timeout in seconds
|
||||
"""
|
||||
self.base_url = (base_url or settings.portainer_url).rstrip("/")
|
||||
self.api_key = api_key or settings.portainer_api_key
|
||||
self.timeout = timeout
|
||||
|
||||
if not self.api_key:
|
||||
logger.warning("Portainer API key not configured")
|
||||
|
||||
def _get_headers(self) -> Dict[str, str]:
|
||||
"""Get request headers with authentication"""
|
||||
return {
|
||||
"X-API-Key": self.api_key,
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""
|
||||
Check if Portainer API is accessible
|
||||
|
||||
Returns:
|
||||
True if accessible, False otherwise
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(f"{self.base_url}/api/status")
|
||||
return response.status_code == 200
|
||||
except Exception as e:
|
||||
logger.error(f"Portainer health check failed: {e}")
|
||||
return False
|
||||
|
||||
async def get_endpoints(self) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all Portainer endpoints (Docker environments)
|
||||
|
||||
Returns:
|
||||
List of endpoint configurations
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/endpoints",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_stacks(self, endpoint_id: Optional[int] = None) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
List all stacks
|
||||
|
||||
Args:
|
||||
endpoint_id: Filter by specific endpoint (optional)
|
||||
|
||||
Returns:
|
||||
List of stack configurations
|
||||
"""
|
||||
params = {}
|
||||
if endpoint_id:
|
||||
params["endpointId"] = endpoint_id
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/stacks",
|
||||
headers=self._get_headers(),
|
||||
params=params
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def get_stack(self, stack_id: int) -> Dict[str, Any]:
|
||||
"""
|
||||
Get details of a specific stack
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
|
||||
Returns:
|
||||
Stack configuration details
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.get(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers()
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def create_stack(
|
||||
self,
|
||||
name: str,
|
||||
stack_file_content: str,
|
||||
endpoint_id: int
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create a new stack from compose file content
|
||||
|
||||
Args:
|
||||
name: Stack name
|
||||
stack_file_content: Docker Compose YAML content
|
||||
endpoint_id: Portainer endpoint to deploy to
|
||||
|
||||
Returns:
|
||||
Created stack details
|
||||
"""
|
||||
payload = {
|
||||
"name": name,
|
||||
"stackFileContent": stack_file_content
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/api/stacks/create/standalone/string",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id},
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def update_stack(
|
||||
self,
|
||||
stack_id: int,
|
||||
stack_file_content: str,
|
||||
endpoint_id: int,
|
||||
prune: bool = False,
|
||||
pull_image: bool = False
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Update an existing stack
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
stack_file_content: New Docker Compose YAML content
|
||||
endpoint_id: Portainer endpoint
|
||||
prune: Remove services no longer defined
|
||||
pull_image: Pull latest images before deployment
|
||||
|
||||
Returns:
|
||||
Updated stack details
|
||||
"""
|
||||
payload = {
|
||||
"stackFileContent": stack_file_content,
|
||||
"prune": prune,
|
||||
"pullImage": pull_image
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.put(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id},
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def delete_stack(self, stack_id: int, endpoint_id: int) -> bool:
|
||||
"""
|
||||
Delete a stack
|
||||
|
||||
Args:
|
||||
stack_id: Stack identifier
|
||||
endpoint_id: Portainer endpoint
|
||||
|
||||
Returns:
|
||||
True if successful
|
||||
"""
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.delete(
|
||||
f"{self.base_url}/api/stacks/{stack_id}",
|
||||
headers=self._get_headers(),
|
||||
params={"endpointId": endpoint_id}
|
||||
)
|
||||
response.raise_for_status()
|
||||
return True
|
||||
|
||||
|
||||
# Singleton instance
|
||||
_portainer_client: Optional[PortainerClient] = None
|
||||
|
||||
|
||||
def get_portainer_client() -> PortainerClient:
|
||||
"""Get singleton Portainer client instance"""
|
||||
global _portainer_client
|
||||
if _portainer_client is None:
|
||||
_portainer_client = PortainerClient()
|
||||
return _portainer_client
|
||||
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
Global configuration for Core Code API
|
||||
"""
|
||||
from pydantic_settings import BaseSettings
|
||||
from functools import lru_cache
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
"""Global application settings"""
|
||||
|
||||
# Application
|
||||
app_name: str = "Core Code API"
|
||||
app_version: str = "1.0.0"
|
||||
debug: bool = False
|
||||
|
||||
# Server
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 8083
|
||||
|
||||
# CORS
|
||||
cors_origins: list[str] = ["*"]
|
||||
cors_credentials: bool = True
|
||||
cors_methods: list[str] = ["*"]
|
||||
cors_headers: list[str] = ["*"]
|
||||
|
||||
# Logging
|
||||
log_level: str = "INFO"
|
||||
|
||||
# Ollama Configuration (for AI orchestration)
|
||||
ollama_base_url: str = "http://ollama:11434"
|
||||
ollama_timeout: int = 300 # 5 minutes
|
||||
|
||||
# Model Configuration
|
||||
default_model: str = "gemma:7b"
|
||||
lightweight_models: str = "gemma:2b,gemma:7b"
|
||||
heavy_models: str = "mistral:7b,gemma2:9b,mixtral:8x7b"
|
||||
code_models: str = "codestral:latest,codegemma:latest"
|
||||
|
||||
# Model Aliases (OpenAI → Local)
|
||||
alias_gpt35: str = "gemma:7b"
|
||||
alias_gpt4: str = "mistral:7b"
|
||||
alias_gpt4_turbo: str = "mixtral:8x7b"
|
||||
alias_gpt4_code: str = "codestral:latest"
|
||||
|
||||
# Memory Configuration
|
||||
memory_tier1_max_turns: int = 10
|
||||
memory_consolidation_threshold: int = 10
|
||||
|
||||
# Qdrant Configuration
|
||||
qdrant_host: str = "qdrant"
|
||||
qdrant_port: int = 6333
|
||||
qdrant_collection_conversations: str = "core_api_conversations"
|
||||
qdrant_collection_documents: str = "core_api_documents"
|
||||
qdrant_collection_user_facts: str = "core_api_user_facts"
|
||||
|
||||
# Embeddings
|
||||
embedding_model: str = "sentence-transformers/all-MiniLM-L6-v2"
|
||||
embedding_dimension: int = 384
|
||||
embedding_batch_size: int = 32
|
||||
|
||||
# Infrastructure Management
|
||||
portainer_url: str = "http://localhost:8001"
|
||||
portainer_api_key: str = ""
|
||||
|
||||
npm_url: str = "http://localhost:81"
|
||||
npm_email: str = ""
|
||||
npm_password: str = ""
|
||||
|
||||
kuma_url: str = "http://localhost:3001"
|
||||
kuma_username: str = ""
|
||||
kuma_password: str = ""
|
||||
|
||||
@property
|
||||
def model_aliases(self) -> dict:
|
||||
"""Computed property for model aliases"""
|
||||
return {
|
||||
"gpt-3.5-turbo": self.alias_gpt35,
|
||||
"gpt-4": self.alias_gpt4,
|
||||
"gpt-4-turbo": self.alias_gpt4_turbo,
|
||||
"gpt-4-code": self.alias_gpt4_code,
|
||||
}
|
||||
|
||||
def get_lightweight_models(self) -> list[str]:
|
||||
"""Parse comma-separated lightweight models"""
|
||||
return [m.strip().strip('"').strip("'") for m in self.lightweight_models.split(",") if m.strip()]
|
||||
|
||||
def get_heavy_models(self) -> list[str]:
|
||||
"""Parse comma-separated heavy models"""
|
||||
return [m.strip().strip('"').strip("'") for m in self.heavy_models.split(",") if m.strip()]
|
||||
|
||||
def get_code_models(self) -> list[str]:
|
||||
"""Parse comma-separated code models"""
|
||||
return [m.strip().strip('"').strip("'") for m in self.code_models.split(",") if m.strip()]
|
||||
|
||||
class Config:
|
||||
env_file = ".env"
|
||||
case_sensitive = False
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def get_settings() -> Settings:
|
||||
"""Cached settings instance"""
|
||||
return Settings()
|
||||
@@ -0,0 +1,5 @@
|
||||
"""
|
||||
Controllers package for Core-API
|
||||
|
||||
Provides controller-based routing architecture for better code organization.
|
||||
"""
|
||||
@@ -0,0 +1,50 @@
|
||||
"""
|
||||
Base controller class for Core-API
|
||||
|
||||
Provides common functionality for all controllers.
|
||||
"""
|
||||
from fastapi import APIRouter
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class BaseController(ABC):
|
||||
"""
|
||||
Base controller class with common functionality
|
||||
|
||||
All controllers should inherit from this class and implement
|
||||
the create_router() method to define their endpoints.
|
||||
"""
|
||||
|
||||
def __init__(self, prefix: str, tags: list[str]):
|
||||
"""
|
||||
Initialize base controller
|
||||
|
||||
Args:
|
||||
prefix: URL prefix for this controller's routes
|
||||
tags: OpenAPI tags for documentation grouping
|
||||
"""
|
||||
self.prefix = prefix
|
||||
self.tags = tags
|
||||
self._router = None
|
||||
|
||||
@abstractmethod
|
||||
def create_router(self) -> APIRouter:
|
||||
"""
|
||||
Create and configure the FastAPI router for this controller
|
||||
|
||||
Returns:
|
||||
Configured APIRouter instance with all endpoints
|
||||
"""
|
||||
pass
|
||||
|
||||
@property
|
||||
def router(self) -> APIRouter:
|
||||
"""
|
||||
Get the router instance, creating it if needed
|
||||
|
||||
Returns:
|
||||
APIRouter instance
|
||||
"""
|
||||
if self._router is None:
|
||||
self._router = self.create_router()
|
||||
return self._router
|
||||
@@ -0,0 +1,270 @@
|
||||
"""
|
||||
Infrastructure Management Controller
|
||||
|
||||
Provides API endpoints for automated infrastructure management,
|
||||
including service deployment, configuration, and monitoring setup.
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from typing import List, Dict, Any, Optional
|
||||
from pydantic import BaseModel
|
||||
|
||||
from src.controllers.base import BaseController
|
||||
from src.clients.portainer_client import get_portainer_client
|
||||
from src.clients.npm_client import get_npm_client
|
||||
from src.logging_config import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
# Response models
|
||||
class ServiceInfo(BaseModel):
|
||||
"""Information about a deployed service"""
|
||||
name: str
|
||||
stack_id: Optional[int]
|
||||
status: str
|
||||
endpoint_id: Optional[int]
|
||||
ports: List[int] = []
|
||||
domains: List[str] = []
|
||||
|
||||
|
||||
class PortInfo(BaseModel):
|
||||
"""Information about an allocated port"""
|
||||
port: int
|
||||
service: str
|
||||
protocol: str = "tcp"
|
||||
description: str = ""
|
||||
|
||||
|
||||
class DomainInfo(BaseModel):
|
||||
"""Information about a configured domain"""
|
||||
domain: str
|
||||
service: str
|
||||
proxy_host_id: Optional[int]
|
||||
ssl_enabled: bool = False
|
||||
certificate_id: Optional[int]
|
||||
|
||||
|
||||
class InfrastructureHealth(BaseModel):
|
||||
"""Overall infrastructure health status"""
|
||||
portainer_connected: bool
|
||||
npm_connected: bool
|
||||
total_stacks: int
|
||||
total_proxy_hosts: int
|
||||
|
||||
|
||||
class InfrastructureController(BaseController):
|
||||
"""
|
||||
Controller for infrastructure management operations
|
||||
|
||||
Provides endpoints for:
|
||||
- Service discovery and listing
|
||||
- Port allocation management
|
||||
- Domain/proxy configuration
|
||||
- Automated service deployment
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(prefix="/infrastructure", tags=["Infrastructure"])
|
||||
|
||||
def create_router(self) -> APIRouter:
|
||||
"""Create and configure the router"""
|
||||
router = APIRouter(prefix=self.prefix, tags=self.tags)
|
||||
|
||||
@router.get(
|
||||
"/health",
|
||||
response_model=InfrastructureHealth,
|
||||
summary="Infrastructure health check"
|
||||
)
|
||||
async def get_infrastructure_health():
|
||||
"""
|
||||
Check health of all infrastructure services
|
||||
|
||||
Returns status of Portainer, NPM, and summary statistics.
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
npm = get_npm_client()
|
||||
|
||||
portainer_healthy = await portainer.health_check()
|
||||
npm_healthy = await npm.health_check()
|
||||
|
||||
total_stacks = 0
|
||||
total_proxy_hosts = 0
|
||||
|
||||
if portainer_healthy:
|
||||
try:
|
||||
stacks = await portainer.get_stacks()
|
||||
total_stacks = len(stacks)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get stacks count: {e}")
|
||||
|
||||
if npm_healthy:
|
||||
try:
|
||||
proxy_hosts = await npm.get_proxy_hosts()
|
||||
total_proxy_hosts = len(proxy_hosts)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get proxy hosts count: {e}")
|
||||
|
||||
return InfrastructureHealth(
|
||||
portainer_connected=portainer_healthy,
|
||||
npm_connected=npm_healthy,
|
||||
total_stacks=total_stacks,
|
||||
total_proxy_hosts=total_proxy_hosts
|
||||
)
|
||||
|
||||
@router.get(
|
||||
"/services",
|
||||
response_model=List[ServiceInfo],
|
||||
summary="List all deployed services"
|
||||
)
|
||||
async def list_services():
|
||||
"""
|
||||
List all deployed services from Portainer stacks
|
||||
|
||||
Returns comprehensive service information including:
|
||||
- Stack/service name
|
||||
- Status
|
||||
- Exposed ports
|
||||
- Configured domains
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
npm = get_npm_client()
|
||||
|
||||
try:
|
||||
stacks = await portainer.get_stacks()
|
||||
proxy_hosts = await npm.get_proxy_hosts()
|
||||
|
||||
# Build domain mapping (domain -> service name)
|
||||
domain_map = {}
|
||||
for proxy in proxy_hosts:
|
||||
for domain in proxy.get("domain_names", []):
|
||||
# Try to extract service name from forward_host
|
||||
forward_host = proxy.get("forward_host", "")
|
||||
domain_map[domain] = forward_host
|
||||
|
||||
services = []
|
||||
for stack in stacks:
|
||||
# Find domains for this stack
|
||||
stack_name = stack.get("Name", "")
|
||||
domains = [
|
||||
domain for domain, host in domain_map.items()
|
||||
if stack_name in host or host in stack_name
|
||||
]
|
||||
|
||||
service_info = ServiceInfo(
|
||||
name=stack_name,
|
||||
stack_id=stack.get("Id"),
|
||||
status=stack.get("Status", "unknown"),
|
||||
endpoint_id=stack.get("EndpointId"),
|
||||
ports=[], # TODO: Extract from stack file
|
||||
domains=domains
|
||||
)
|
||||
services.append(service_info)
|
||||
|
||||
return services
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list services: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@router.get(
|
||||
"/services/{name}",
|
||||
response_model=ServiceInfo,
|
||||
summary="Get service details"
|
||||
)
|
||||
async def get_service(name: str):
|
||||
"""
|
||||
Get detailed information about a specific service
|
||||
|
||||
Args:
|
||||
name: Service/stack name
|
||||
|
||||
Returns:
|
||||
Service details including status and configuration
|
||||
"""
|
||||
portainer = get_portainer_client()
|
||||
|
||||
try:
|
||||
stacks = await portainer.get_stacks()
|
||||
|
||||
# Find stack by name (case-insensitive)
|
||||
stack = next(
|
||||
(s for s in stacks if s.get("Name", "").lower() == name.lower()),
|
||||
None
|
||||
)
|
||||
|
||||
if not stack:
|
||||
raise HTTPException(status_code=404, detail=f"Service '{name}' not found")
|
||||
|
||||
return ServiceInfo(
|
||||
name=stack.get("Name", ""),
|
||||
stack_id=stack.get("Id"),
|
||||
status=stack.get("Status", "unknown"),
|
||||
endpoint_id=stack.get("EndpointId"),
|
||||
ports=[],
|
||||
domains=[]
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to get service '{name}': {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@router.get(
|
||||
"/ports",
|
||||
response_model=List[PortInfo],
|
||||
summary="List allocated ports"
|
||||
)
|
||||
async def list_ports():
|
||||
"""
|
||||
List all currently allocated ports
|
||||
|
||||
Scans services and proxy configurations to build
|
||||
a comprehensive port allocation map.
|
||||
"""
|
||||
# TODO: Implement port scanning from containers and proxy configs
|
||||
# For now, return a placeholder
|
||||
return []
|
||||
|
||||
@router.get(
|
||||
"/domains",
|
||||
response_model=List[DomainInfo],
|
||||
summary="List configured domains"
|
||||
)
|
||||
async def list_domains():
|
||||
"""
|
||||
List all configured domain names
|
||||
|
||||
Returns domain-to-service mappings with SSL status.
|
||||
"""
|
||||
npm = get_npm_client()
|
||||
|
||||
try:
|
||||
proxy_hosts = await npm.get_proxy_hosts()
|
||||
|
||||
domains = []
|
||||
for proxy in proxy_hosts:
|
||||
service_name = proxy.get("forward_host", "localhost")
|
||||
certificate_id = proxy.get("certificate_id", 0)
|
||||
|
||||
for domain in proxy.get("domain_names", []):
|
||||
domain_info = DomainInfo(
|
||||
domain=domain,
|
||||
service=service_name,
|
||||
proxy_host_id=proxy.get("id"),
|
||||
ssl_enabled=certificate_id > 0,
|
||||
certificate_id=certificate_id if certificate_id > 0 else None
|
||||
)
|
||||
domains.append(domain_info)
|
||||
|
||||
return domains
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to list domains: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
return router
|
||||
|
||||
|
||||
# Create controller instance
|
||||
infrastructure_controller = InfrastructureController()
|
||||
@@ -0,0 +1,24 @@
|
||||
"""
|
||||
Infrastructure Credentials Template
|
||||
|
||||
INSTRUCTIONS:
|
||||
1. Copy this file to credentials.py
|
||||
2. Fill in your actual credentials
|
||||
3. DO NOT commit credentials.py to version control (it's in .gitignore)
|
||||
|
||||
This file should be committed to the repository as a template.
|
||||
"""
|
||||
|
||||
# Portainer Configuration
|
||||
PORTAINER_URL = "http://localhost:8001"
|
||||
PORTAINER_API_KEY = "ptr_your_api_token_here" # Create in Portainer UI: User menu → My account → Access tokens
|
||||
|
||||
# Nginx Proxy Manager Configuration
|
||||
NPM_URL = "http://localhost:81"
|
||||
NPM_EMAIL = "admin@example.com"
|
||||
NPM_PASSWORD = "your_password_here"
|
||||
|
||||
# Uptime Kuma Configuration
|
||||
KUMA_URL = "http://localhost:3001"
|
||||
KUMA_USERNAME = "admin"
|
||||
KUMA_PASSWORD = "your_password_here"
|
||||
@@ -0,0 +1,49 @@
|
||||
"""
|
||||
Logging configuration for Core Code API
|
||||
"""
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def setup_logging(log_level: str = "INFO") -> None:
|
||||
"""
|
||||
Configure logging for the application
|
||||
|
||||
Args:
|
||||
log_level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL)
|
||||
"""
|
||||
# Create logs directory if it doesn't exist
|
||||
log_dir = Path("logs")
|
||||
log_dir.mkdir(exist_ok=True)
|
||||
|
||||
# Configure root logger
|
||||
logging.basicConfig(
|
||||
level=getattr(logging, log_level.upper()),
|
||||
format="%(asctime)s | %(levelname)-8s | %(name)s:%(funcName)s:%(lineno)d | %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
handlers=[
|
||||
# Console handler
|
||||
logging.StreamHandler(sys.stdout),
|
||||
# File handler
|
||||
logging.FileHandler(log_dir / "app.log", encoding="utf-8")
|
||||
]
|
||||
)
|
||||
|
||||
# Set specific log levels for third-party libraries
|
||||
logging.getLogger("httpx").setLevel(logging.WARNING)
|
||||
logging.getLogger("httpcore").setLevel(logging.WARNING)
|
||||
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||
|
||||
|
||||
def get_logger(name: str) -> logging.Logger:
|
||||
"""
|
||||
Get a logger instance
|
||||
|
||||
Args:
|
||||
name: Logger name (typically __name__)
|
||||
|
||||
Returns:
|
||||
Configured logger instance
|
||||
"""
|
||||
return logging.getLogger(name)
|
||||
@@ -0,0 +1,200 @@
|
||||
"""
|
||||
Main FastAPI application for Core Code API
|
||||
"""
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from src.config import get_settings
|
||||
from src.logging_config import setup_logging, get_logger
|
||||
from src.web_scraper import router as web_scraper_router
|
||||
from src.api.v1.chat import router as chat_router
|
||||
from src.api.v1.models import router as models_router
|
||||
from src.api.v1.conversations import router as conversations_router
|
||||
from src.models.ollama_client import get_ollama_client, close_ollama_client
|
||||
|
||||
# Initialize settings
|
||||
settings = get_settings()
|
||||
|
||||
# Setup logging
|
||||
setup_logging(settings.log_level)
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
"""
|
||||
Application lifespan manager for startup/shutdown events
|
||||
|
||||
Args:
|
||||
app: FastAPI application instance
|
||||
"""
|
||||
# Startup
|
||||
logger.info("=" * 60)
|
||||
logger.info(f"Starting {settings.app_name} v{settings.app_version}")
|
||||
logger.info(f"Debug mode: {settings.debug}")
|
||||
logger.info(f"Log level: {settings.log_level}")
|
||||
logger.info(f"Ollama URL: {settings.ollama_base_url}")
|
||||
logger.info("=" * 60)
|
||||
|
||||
# Check Ollama connectivity
|
||||
ollama_client = get_ollama_client()
|
||||
ollama_healthy = await ollama_client.health_check()
|
||||
if ollama_healthy:
|
||||
logger.info("✓ Ollama connection successful")
|
||||
else:
|
||||
logger.warning("✗ Ollama connection failed - AI features may not work")
|
||||
|
||||
yield
|
||||
|
||||
# Shutdown
|
||||
logger.info("Shutting down application")
|
||||
await close_ollama_client()
|
||||
|
||||
|
||||
# Create FastAPI application
|
||||
app = FastAPI(
|
||||
title=settings.app_name,
|
||||
version=settings.app_version,
|
||||
description="""
|
||||
Core Code API provides OpenAPI-compatible functions and AI orchestration for Open WebUI.
|
||||
|
||||
## Features
|
||||
|
||||
### OpenAI-Compatible API (v1)
|
||||
- `/v1/chat/completions` - Chat completions with streaming support
|
||||
- `/v1/models` - List available models
|
||||
Compatible with OpenAI client libraries and Open WebUI.
|
||||
|
||||
### Conversation Memory (Phase 2)
|
||||
- `/v1/conversations/{id}` - Get conversation history
|
||||
- `/v1/conversations/{id}/search` - Semantic search within conversation
|
||||
- `/v1/conversations/search` - Search across all conversations
|
||||
- `/v1/conversations/{id}/stats` - Get conversation statistics
|
||||
- `/v1/conversations/{id}/consolidate` - Manual consolidation
|
||||
- `DELETE /v1/conversations/{id}` - Delete conversation
|
||||
|
||||
Multi-tier memory system:
|
||||
- **Tier 1**: Fast in-memory buffer (last 10 turns)
|
||||
- **Tier 2/3**: Unified Qdrant storage (persistent + semantic search)
|
||||
|
||||
### Web Scraper
|
||||
Intelligent web scraping with main content extraction.
|
||||
Perfect for extracting articles, documentation, and blog posts for LLM consumption.
|
||||
|
||||
## Integration
|
||||
|
||||
This API is designed to integrate with:
|
||||
- **Open WebUI**: Direct OpenAI API compatibility
|
||||
- **Open WebUI Functions**: Import via OpenAPI spec
|
||||
- **Open WebUI Pipelines**: Use as data source
|
||||
- **LangChain**: Compatible with standard HTTP tools
|
||||
|
||||
## Documentation
|
||||
|
||||
- **OpenAPI Spec**: `/openapi.json`
|
||||
- **Swagger UI**: `/docs`
|
||||
- **ReDoc**: `/redoc`
|
||||
""",
|
||||
docs_url="/docs",
|
||||
redoc_url="/redoc",
|
||||
openapi_url="/openapi.json",
|
||||
lifespan=lifespan,
|
||||
debug=settings.debug
|
||||
)
|
||||
|
||||
# Add CORS middleware
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=settings.cors_origins,
|
||||
allow_credentials=settings.cors_credentials,
|
||||
allow_methods=settings.cors_methods,
|
||||
allow_headers=settings.cors_headers,
|
||||
)
|
||||
|
||||
|
||||
# Root endpoint
|
||||
@app.get(
|
||||
"/",
|
||||
tags=["Health"],
|
||||
summary="Service information",
|
||||
response_class=JSONResponse
|
||||
)
|
||||
async def root():
|
||||
"""
|
||||
Get service information and health status
|
||||
|
||||
Returns basic information about the API service and available endpoints.
|
||||
"""
|
||||
logger.debug("Root endpoint accessed")
|
||||
return {
|
||||
"service": settings.app_name,
|
||||
"version": settings.app_version,
|
||||
"status": "healthy",
|
||||
"documentation": {
|
||||
"swagger_ui": "/docs",
|
||||
"redoc": "/redoc",
|
||||
"openapi_spec": "/openapi.json"
|
||||
},
|
||||
"endpoints": {
|
||||
"chat_completions": "/v1/chat/completions",
|
||||
"models": "/v1/models",
|
||||
"conversations": "/v1/conversations",
|
||||
"web_scraper": "/web-scraper/scrape",
|
||||
"health": "/health"
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
# Health check endpoint
|
||||
@app.get(
|
||||
"/health",
|
||||
tags=["Health"],
|
||||
summary="Health check",
|
||||
response_class=JSONResponse
|
||||
)
|
||||
async def health_check():
|
||||
"""
|
||||
Simple health check endpoint for container orchestration
|
||||
|
||||
Returns a 200 OK status when the service is running properly.
|
||||
Used by Docker, Kubernetes, and load balancers.
|
||||
"""
|
||||
ollama_client = get_ollama_client()
|
||||
ollama_healthy = await ollama_client.health_check()
|
||||
|
||||
return {
|
||||
"status": "healthy",
|
||||
"ollama_connected": ollama_healthy
|
||||
}
|
||||
|
||||
|
||||
# Include routers
|
||||
app.include_router(chat_router) # /v1/chat/completions
|
||||
app.include_router(models_router) # /v1/models
|
||||
app.include_router(conversations_router) # /v1/conversations
|
||||
app.include_router(web_scraper_router) # /web-scraper/scrape
|
||||
|
||||
|
||||
# Global exception handler
|
||||
@app.exception_handler(Exception)
|
||||
async def global_exception_handler(request, exc):
|
||||
"""
|
||||
Catch-all exception handler for unhandled errors
|
||||
|
||||
Args:
|
||||
request: The request that caused the exception
|
||||
exc: The exception instance
|
||||
|
||||
Returns:
|
||||
JSON error response
|
||||
"""
|
||||
logger.error(f"Unhandled exception: {str(exc)}", exc_info=True)
|
||||
return JSONResponse(
|
||||
status_code=500,
|
||||
content={
|
||||
"detail": "Internal server error",
|
||||
"type": type(exc).__name__
|
||||
}
|
||||
)
|
||||
@@ -0,0 +1,42 @@
|
||||
"""
|
||||
Memory system for conversation persistence
|
||||
|
||||
Simplified architecture:
|
||||
- Tier 1: ConversationBufferMemory (in-memory, fast, last 10 turns)
|
||||
- Tier 2/3: QdrantConversationMemory (unified persistent + semantic search)
|
||||
- Manager: MemoryManager (orchestrates all tiers)
|
||||
"""
|
||||
from .tier1_buffer import ConversationBufferMemory, get_buffer_memory
|
||||
from .qdrant_memory import QdrantConversationMemory, get_qdrant_memory
|
||||
from .manager import MemoryManager, get_memory_manager
|
||||
from .schemas import (
|
||||
ConversationTurn,
|
||||
ConversationBuffer,
|
||||
ConversationMetadata,
|
||||
ConversationSummary,
|
||||
MemoryQuery,
|
||||
MemoryResult,
|
||||
MessageRole,
|
||||
TokenUsage
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# Manager (primary interface)
|
||||
"MemoryManager",
|
||||
"get_memory_manager",
|
||||
# Tier 1
|
||||
"ConversationBufferMemory",
|
||||
"get_buffer_memory",
|
||||
# Tier 2/3
|
||||
"QdrantConversationMemory",
|
||||
"get_qdrant_memory",
|
||||
# Schemas
|
||||
"ConversationTurn",
|
||||
"ConversationBuffer",
|
||||
"ConversationMetadata",
|
||||
"ConversationSummary",
|
||||
"MemoryQuery",
|
||||
"MemoryResult",
|
||||
"MessageRole",
|
||||
"TokenUsage",
|
||||
]
|
||||
@@ -0,0 +1,169 @@
|
||||
"""
|
||||
Base classes for memory system
|
||||
"""
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional
|
||||
from .schemas import ConversationTurn, ConversationBuffer, MemoryQuery, MemoryResult
|
||||
|
||||
|
||||
class BaseMemory(ABC):
|
||||
"""Base class for all memory tiers"""
|
||||
|
||||
@abstractmethod
|
||||
async def add_turn(self, conversation_id: str, turn: ConversationTurn) -> None:
|
||||
"""
|
||||
Add a new turn to memory
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
turn: The conversation turn to store
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_turns(
|
||||
self,
|
||||
conversation_id: str,
|
||||
limit: Optional[int] = None,
|
||||
offset: int = 0
|
||||
) -> List[ConversationTurn]:
|
||||
"""
|
||||
Retrieve turns from memory
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
limit: Maximum number of turns to retrieve
|
||||
offset: Number of turns to skip
|
||||
|
||||
Returns:
|
||||
List of conversation turns
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def clear_conversation(self, conversation_id: str) -> None:
|
||||
"""
|
||||
Clear all turns for a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def conversation_exists(self, conversation_id: str) -> bool:
|
||||
"""
|
||||
Check if a conversation exists in this memory tier
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
True if conversation exists
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class Tier1Memory(BaseMemory):
|
||||
"""Base class for Tier 1 (working memory)"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_buffer(self, conversation_id: str) -> Optional[ConversationBuffer]:
|
||||
"""
|
||||
Get the full conversation buffer
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
ConversationBuffer or None if not found
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def prune(self, conversation_id: str, keep_last: int = 5) -> None:
|
||||
"""
|
||||
Prune old turns, keeping only the most recent ones
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
keep_last: Number of recent turns to keep
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class Tier2Memory(BaseMemory):
|
||||
"""Base class for Tier 2 (short-term memory with summaries)"""
|
||||
|
||||
@abstractmethod
|
||||
async def add_summary(
|
||||
self,
|
||||
conversation_id: str,
|
||||
summary_text: str,
|
||||
turn_range_start: int,
|
||||
turn_range_end: int
|
||||
) -> None:
|
||||
"""
|
||||
Add a conversation summary
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
summary_text: The summarized text
|
||||
turn_range_start: First turn number in summary
|
||||
turn_range_end: Last turn number in summary
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_summaries(self, conversation_id: str) -> List[dict]:
|
||||
"""
|
||||
Get all summaries for a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
List of summary dictionaries
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
class Tier3Memory(BaseMemory):
|
||||
"""Base class for Tier 3 (long-term vector memory)"""
|
||||
|
||||
@abstractmethod
|
||||
async def add_turn_with_embedding(
|
||||
self,
|
||||
conversation_id: str,
|
||||
turn: ConversationTurn,
|
||||
embedding: List[float]
|
||||
) -> None:
|
||||
"""
|
||||
Add a turn with its vector embedding
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
turn: The conversation turn
|
||||
embedding: Vector embedding of the turn content
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def similarity_search(
|
||||
self,
|
||||
query_embedding: List[float],
|
||||
conversation_id: Optional[str] = None,
|
||||
limit: int = 5
|
||||
) -> List[dict]:
|
||||
"""
|
||||
Perform semantic similarity search
|
||||
|
||||
Args:
|
||||
query_embedding: Vector embedding of the search query
|
||||
conversation_id: Optional filter to specific conversation
|
||||
limit: Maximum number of results
|
||||
|
||||
Returns:
|
||||
List of matching turns with scores
|
||||
"""
|
||||
pass
|
||||
@@ -0,0 +1,319 @@
|
||||
"""
|
||||
Memory Manager: Orchestrates all memory tiers
|
||||
|
||||
Coordinates:
|
||||
- Tier 1: ConversationBufferMemory (RAM, fast, last N turns)
|
||||
- Tier 2/3: QdrantConversationMemory (persistent + semantic)
|
||||
|
||||
Provides unified interface for memory operations with automatic
|
||||
tier management and consolidation.
|
||||
"""
|
||||
import logging
|
||||
import asyncio
|
||||
from typing import List, Optional, Dict, Any
|
||||
from datetime import datetime
|
||||
|
||||
from .tier1_buffer import ConversationBufferMemory, get_buffer_memory
|
||||
from .qdrant_memory import QdrantConversationMemory, get_qdrant_memory
|
||||
from .schemas import ConversationTurn, MessageRole, TokenUsage
|
||||
from src.config import get_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class MemoryManager:
|
||||
"""
|
||||
Unified memory manager orchestrating all tiers
|
||||
|
||||
Responsibilities:
|
||||
- Add turns to appropriate tiers
|
||||
- Retrieve conversation history (buffer + persistent)
|
||||
- Consolidate buffer to persistent storage
|
||||
- Semantic search across all conversations
|
||||
- Memory lifecycle management
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
buffer_memory: Optional[ConversationBufferMemory] = None,
|
||||
qdrant_memory: Optional[QdrantConversationMemory] = None,
|
||||
auto_consolidate: bool = True
|
||||
):
|
||||
"""
|
||||
Initialize memory manager
|
||||
|
||||
Args:
|
||||
buffer_memory: Optional Tier 1 buffer instance
|
||||
qdrant_memory: Optional Tier 2/3 Qdrant instance
|
||||
auto_consolidate: Automatically consolidate when buffer threshold reached
|
||||
"""
|
||||
self.buffer_memory = buffer_memory or get_buffer_memory()
|
||||
self.qdrant_memory = qdrant_memory or get_qdrant_memory()
|
||||
self.auto_consolidate = auto_consolidate
|
||||
|
||||
logger.info(
|
||||
f"MemoryManager initialized (auto_consolidate={auto_consolidate})"
|
||||
)
|
||||
|
||||
async def add_turn(
|
||||
self,
|
||||
conversation_id: str,
|
||||
role: MessageRole,
|
||||
content: str,
|
||||
tokens: Optional[TokenUsage] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
) -> ConversationTurn:
|
||||
"""
|
||||
Add a conversation turn to memory
|
||||
|
||||
Automatically:
|
||||
1. Adds to Tier 1 (buffer)
|
||||
2. Checks if consolidation threshold reached
|
||||
3. Consolidates to Tier 2/3 if needed
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
role: Message role (user, assistant, system)
|
||||
content: Message content
|
||||
tokens: Optional token usage
|
||||
metadata: Optional metadata
|
||||
|
||||
Returns:
|
||||
The created conversation turn
|
||||
"""
|
||||
# Get current buffer to determine turn number
|
||||
buffer = await self.buffer_memory.get_buffer(conversation_id)
|
||||
turn_number = (buffer.metadata.turn_count + 1) if buffer else 1
|
||||
|
||||
# Create turn
|
||||
turn = ConversationTurn(
|
||||
role=role,
|
||||
content=content,
|
||||
timestamp=datetime.utcnow(),
|
||||
turn_number=turn_number,
|
||||
tokens=tokens,
|
||||
metadata=metadata or {}
|
||||
)
|
||||
|
||||
# Add to Tier 1 (buffer)
|
||||
await self.buffer_memory.add_turn(conversation_id, turn)
|
||||
logger.debug(f"Turn {turn_number} added to buffer for {conversation_id}")
|
||||
|
||||
# Check consolidation threshold
|
||||
if self.auto_consolidate:
|
||||
buffer = await self.buffer_memory.get_buffer(conversation_id)
|
||||
if buffer.metadata.turn_count >= settings.memory_consolidation_threshold:
|
||||
logger.info(
|
||||
f"Consolidation threshold reached for {conversation_id} "
|
||||
f"({buffer.metadata.turn_count} turns)"
|
||||
)
|
||||
await self._consolidate_buffer(conversation_id)
|
||||
|
||||
return turn
|
||||
|
||||
async def get_recent_turns(
|
||||
self,
|
||||
conversation_id: str,
|
||||
limit: int = 10
|
||||
) -> List[ConversationTurn]:
|
||||
"""
|
||||
Get recent conversation turns (from buffer)
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
limit: Maximum number of turns to retrieve
|
||||
|
||||
Returns:
|
||||
List of recent conversation turns
|
||||
"""
|
||||
return await self.buffer_memory.get_recent_turns(conversation_id, limit)
|
||||
|
||||
async def get_full_history(
|
||||
self,
|
||||
conversation_id: str,
|
||||
include_buffer: bool = True
|
||||
) -> List[ConversationTurn]:
|
||||
"""
|
||||
Get complete conversation history
|
||||
|
||||
Combines:
|
||||
- Tier 2/3: Persistent history from Qdrant
|
||||
- Tier 1: Recent buffer (if include_buffer=True)
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
include_buffer: Include buffer turns not yet consolidated
|
||||
|
||||
Returns:
|
||||
Complete conversation history, sorted chronologically
|
||||
"""
|
||||
# Get from Qdrant (Tier 2)
|
||||
qdrant_turns = await self.qdrant_memory.get_turns(conversation_id)
|
||||
|
||||
# Get from buffer (Tier 1)
|
||||
if include_buffer:
|
||||
buffer_turns = await self.buffer_memory.get_turns(conversation_id)
|
||||
|
||||
# Combine and deduplicate (Qdrant is source of truth)
|
||||
qdrant_turn_numbers = {t.turn_number for t in qdrant_turns}
|
||||
new_buffer_turns = [
|
||||
t for t in buffer_turns
|
||||
if t.turn_number not in qdrant_turn_numbers
|
||||
]
|
||||
|
||||
all_turns = qdrant_turns + new_buffer_turns
|
||||
else:
|
||||
all_turns = qdrant_turns
|
||||
|
||||
# Sort chronologically
|
||||
all_turns.sort(key=lambda t: t.turn_number)
|
||||
|
||||
return all_turns
|
||||
|
||||
async def search_conversations(
|
||||
self,
|
||||
query: str,
|
||||
conversation_id: Optional[str] = None,
|
||||
limit: int = 5
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Semantic search across conversations (Tier 3 mode)
|
||||
|
||||
Args:
|
||||
query: Search query
|
||||
conversation_id: Optional filter to specific conversation
|
||||
limit: Maximum number of results
|
||||
|
||||
Returns:
|
||||
List of matching turns with scores
|
||||
"""
|
||||
return await self.qdrant_memory.similarity_search(
|
||||
query=query,
|
||||
conversation_id=conversation_id,
|
||||
limit=limit
|
||||
)
|
||||
|
||||
async def consolidate(self, conversation_id: str) -> int:
|
||||
"""
|
||||
Manually trigger consolidation for a conversation
|
||||
|
||||
Moves all buffer turns to Qdrant (Tier 1 → Tier 2/3)
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
Number of turns consolidated
|
||||
"""
|
||||
return await self._consolidate_buffer(conversation_id)
|
||||
|
||||
async def _consolidate_buffer(self, conversation_id: str) -> int:
|
||||
"""
|
||||
Internal consolidation: Move buffer turns to Qdrant
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
Number of turns consolidated
|
||||
"""
|
||||
buffer = await self.buffer_memory.get_buffer(conversation_id)
|
||||
if not buffer or len(buffer.turns) == 0:
|
||||
logger.debug(f"No turns to consolidate for {conversation_id}")
|
||||
return 0
|
||||
|
||||
# Get turns from buffer
|
||||
buffer_turns = buffer.turns.copy()
|
||||
|
||||
# Add to Qdrant
|
||||
consolidated_count = 0
|
||||
for turn in buffer_turns:
|
||||
try:
|
||||
await self.qdrant_memory.add_turn(conversation_id, turn)
|
||||
consolidated_count += 1
|
||||
except Exception as e:
|
||||
logger.error(f"Error consolidating turn {turn.turn_number}: {e}")
|
||||
|
||||
logger.info(
|
||||
f"Consolidated {consolidated_count}/{len(buffer_turns)} turns "
|
||||
f"for {conversation_id}"
|
||||
)
|
||||
|
||||
# Note: We keep the buffer, just stored in Qdrant as well
|
||||
# Buffer will be pruned naturally as new turns come in
|
||||
# This provides redundancy and fast access to recent turns
|
||||
|
||||
return consolidated_count
|
||||
|
||||
async def clear_conversation(
|
||||
self,
|
||||
conversation_id: str,
|
||||
clear_buffer: bool = True,
|
||||
clear_qdrant: bool = True
|
||||
) -> None:
|
||||
"""
|
||||
Clear conversation from memory
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
clear_buffer: Clear from Tier 1 buffer
|
||||
clear_qdrant: Clear from Tier 2/3 Qdrant
|
||||
"""
|
||||
if clear_buffer:
|
||||
await self.buffer_memory.clear_conversation(conversation_id)
|
||||
logger.info(f"Cleared buffer for {conversation_id}")
|
||||
|
||||
if clear_qdrant:
|
||||
await self.qdrant_memory.clear_conversation(conversation_id)
|
||||
logger.info(f"Cleared Qdrant for {conversation_id}")
|
||||
|
||||
async def get_conversation_stats(
|
||||
self,
|
||||
conversation_id: str
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Get conversation statistics across all tiers
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
Dictionary with stats from buffer and Qdrant
|
||||
"""
|
||||
# Get buffer stats
|
||||
buffer = await self.buffer_memory.get_buffer(conversation_id)
|
||||
buffer_stats = {
|
||||
"buffer_turns": buffer.metadata.turn_count if buffer else 0,
|
||||
"buffer_tokens": buffer.metadata.total_tokens if buffer else 0
|
||||
}
|
||||
|
||||
# Get Qdrant stats
|
||||
qdrant_stats = await self.qdrant_memory.get_conversation_stats(conversation_id)
|
||||
|
||||
# Combine
|
||||
return {
|
||||
"conversation_id": conversation_id,
|
||||
**buffer_stats,
|
||||
"qdrant_turns": qdrant_stats["total_turns"],
|
||||
"qdrant_tokens": qdrant_stats["total_tokens"],
|
||||
"exists_in_buffer": buffer is not None,
|
||||
"exists_in_qdrant": qdrant_stats["exists"]
|
||||
}
|
||||
|
||||
|
||||
# Global instance
|
||||
_memory_manager: Optional[MemoryManager] = None
|
||||
|
||||
|
||||
def get_memory_manager() -> MemoryManager:
|
||||
"""
|
||||
Get or create global memory manager instance
|
||||
|
||||
Returns:
|
||||
MemoryManager instance
|
||||
"""
|
||||
global _memory_manager
|
||||
if _memory_manager is None:
|
||||
_memory_manager = MemoryManager()
|
||||
return _memory_manager
|
||||
@@ -0,0 +1,387 @@
|
||||
"""
|
||||
Unified Tier 2/3: Qdrant-based conversation memory
|
||||
|
||||
Single Qdrant collection serving both purposes:
|
||||
- Tier 2: Historical retrieval (filter by conversation_id, time-based)
|
||||
- Tier 3: Semantic search (vector similarity across conversations)
|
||||
"""
|
||||
import logging
|
||||
import uuid
|
||||
from typing import List, Optional, Dict, Any
|
||||
from datetime import datetime
|
||||
from qdrant_client import QdrantClient
|
||||
from qdrant_client.models import (
|
||||
Distance,
|
||||
VectorParams,
|
||||
PointStruct,
|
||||
Filter,
|
||||
FieldCondition,
|
||||
MatchValue,
|
||||
Range,
|
||||
)
|
||||
|
||||
from .base import BaseMemory
|
||||
from .schemas import ConversationTurn, MessageRole
|
||||
from src.config import get_settings
|
||||
from src.models.embeddings import get_embedding_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class QdrantConversationMemory(BaseMemory):
|
||||
"""
|
||||
Unified conversation memory using Qdrant
|
||||
|
||||
Stores all conversation turns with vectors for semantic search.
|
||||
Can be queried in two ways:
|
||||
- Tier 2 mode: Filter by conversation_id for chronological history
|
||||
- Tier 3 mode: Vector similarity search for semantic recall
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
collection_name: Optional[str] = None,
|
||||
host: Optional[str] = None,
|
||||
port: Optional[int] = None
|
||||
):
|
||||
"""
|
||||
Initialize Qdrant memory
|
||||
|
||||
Args:
|
||||
collection_name: Name of Qdrant collection
|
||||
host: Qdrant host
|
||||
port: Qdrant port
|
||||
"""
|
||||
self.collection_name = collection_name or settings.qdrant_collection_conversations
|
||||
self.host = host or settings.qdrant_host
|
||||
self.port = port or settings.qdrant_port
|
||||
|
||||
# Initialize clients
|
||||
self.client = QdrantClient(host=self.host, port=self.port)
|
||||
self.embedding_client = get_embedding_client()
|
||||
|
||||
logger.info(
|
||||
f"Initialized QdrantConversationMemory: "
|
||||
f"{self.host}:{self.port}/{self.collection_name}"
|
||||
)
|
||||
|
||||
# Ensure collection exists
|
||||
self._ensure_collection()
|
||||
|
||||
def _ensure_collection(self) -> None:
|
||||
"""Create collection if it doesn't exist"""
|
||||
try:
|
||||
collections = self.client.get_collections().collections
|
||||
collection_names = [c.name for c in collections]
|
||||
|
||||
if self.collection_name not in collection_names:
|
||||
logger.info(f"Creating collection: {self.collection_name}")
|
||||
self.client.create_collection(
|
||||
collection_name=self.collection_name,
|
||||
vectors_config=VectorParams(
|
||||
size=settings.embedding_dimension,
|
||||
distance=Distance.COSINE
|
||||
)
|
||||
)
|
||||
logger.info(f"✓ Collection created: {self.collection_name}")
|
||||
else:
|
||||
logger.info(f"✓ Collection exists: {self.collection_name}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error ensuring collection: {e}")
|
||||
raise
|
||||
|
||||
async def add_turn(self, conversation_id: str, turn: ConversationTurn) -> None:
|
||||
"""
|
||||
Add a conversation turn with its embedding
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
turn: The conversation turn to store
|
||||
"""
|
||||
# Generate embedding
|
||||
embedding = self.embedding_client.embed_text(turn.content)
|
||||
|
||||
# Create point ID: deterministic UUID from conversation_id + turn_number
|
||||
# Qdrant requires UUID or unsigned int, so we generate UUID from string
|
||||
point_id_str = f"{conversation_id}_{turn.turn_number}"
|
||||
point_id = str(uuid.uuid5(uuid.NAMESPACE_DNS, point_id_str))
|
||||
|
||||
# Build payload
|
||||
payload = {
|
||||
"conversation_id": conversation_id,
|
||||
"turn_number": turn.turn_number,
|
||||
"role": turn.role.value if isinstance(turn.role, MessageRole) else turn.role,
|
||||
"content": turn.content,
|
||||
"timestamp": turn.timestamp.isoformat(),
|
||||
"metadata": turn.metadata,
|
||||
}
|
||||
|
||||
# Add token info if available
|
||||
if turn.tokens:
|
||||
payload["tokens_prompt"] = turn.tokens.prompt
|
||||
payload["tokens_completion"] = turn.tokens.completion
|
||||
payload["tokens_total"] = turn.tokens.total
|
||||
|
||||
# Upsert to Qdrant
|
||||
try:
|
||||
self.client.upsert(
|
||||
collection_name=self.collection_name,
|
||||
points=[
|
||||
PointStruct(
|
||||
id=point_id,
|
||||
vector=embedding,
|
||||
payload=payload
|
||||
)
|
||||
]
|
||||
)
|
||||
logger.debug(f"Stored turn {turn.turn_number} for conversation {conversation_id}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error storing turn in Qdrant: {e}")
|
||||
raise
|
||||
|
||||
async def get_turns(
|
||||
self,
|
||||
conversation_id: str,
|
||||
limit: Optional[int] = None,
|
||||
offset: int = 0
|
||||
) -> List[ConversationTurn]:
|
||||
"""
|
||||
Retrieve turns for a conversation (Tier 2 mode: chronological)
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
limit: Maximum number of turns to retrieve
|
||||
offset: Number of turns to skip
|
||||
|
||||
Returns:
|
||||
List of conversation turns
|
||||
"""
|
||||
try:
|
||||
# Scroll through all points for this conversation
|
||||
points, _ = self.client.scroll(
|
||||
collection_name=self.collection_name,
|
||||
scroll_filter=Filter(
|
||||
must=[
|
||||
FieldCondition(
|
||||
key="conversation_id",
|
||||
match=MatchValue(value=conversation_id)
|
||||
)
|
||||
]
|
||||
),
|
||||
limit=limit or 100,
|
||||
offset=offset,
|
||||
with_payload=True,
|
||||
with_vectors=False
|
||||
)
|
||||
|
||||
# Convert to ConversationTurn objects
|
||||
turns = []
|
||||
for point in points:
|
||||
payload = point.payload
|
||||
turn = ConversationTurn(
|
||||
role=MessageRole(payload["role"]),
|
||||
content=payload["content"],
|
||||
timestamp=datetime.fromisoformat(payload["timestamp"]),
|
||||
turn_number=payload["turn_number"],
|
||||
metadata=payload.get("metadata", {})
|
||||
)
|
||||
turns.append(turn)
|
||||
|
||||
# Sort by turn_number
|
||||
turns.sort(key=lambda t: t.turn_number)
|
||||
|
||||
return turns
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error retrieving turns from Qdrant: {e}")
|
||||
return []
|
||||
|
||||
async def similarity_search(
|
||||
self,
|
||||
query: str,
|
||||
conversation_id: Optional[str] = None,
|
||||
limit: int = 5
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Semantic search for relevant turns (Tier 3 mode: semantic)
|
||||
|
||||
Args:
|
||||
query: Search query text
|
||||
conversation_id: Optional filter to specific conversation
|
||||
limit: Maximum number of results
|
||||
|
||||
Returns:
|
||||
List of matching turns with scores
|
||||
"""
|
||||
try:
|
||||
# Generate query embedding
|
||||
query_embedding = self.embedding_client.embed_text(query)
|
||||
|
||||
# Build filter if conversation_id specified
|
||||
search_filter = None
|
||||
if conversation_id:
|
||||
search_filter = Filter(
|
||||
must=[
|
||||
FieldCondition(
|
||||
key="conversation_id",
|
||||
match=MatchValue(value=conversation_id)
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
# Search in Qdrant
|
||||
results = self.client.search(
|
||||
collection_name=self.collection_name,
|
||||
query_vector=query_embedding,
|
||||
query_filter=search_filter,
|
||||
limit=limit,
|
||||
with_payload=True
|
||||
)
|
||||
|
||||
# Convert results
|
||||
matches = []
|
||||
for result in results:
|
||||
payload = result.payload
|
||||
match = {
|
||||
"conversation_id": payload["conversation_id"],
|
||||
"turn_number": payload["turn_number"],
|
||||
"role": payload["role"],
|
||||
"content": payload["content"],
|
||||
"timestamp": payload["timestamp"],
|
||||
"score": result.score,
|
||||
}
|
||||
matches.append(match)
|
||||
|
||||
logger.debug(
|
||||
f"Semantic search found {len(matches)} matches for query: {query[:50]}..."
|
||||
)
|
||||
|
||||
return matches
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in semantic search: {e}")
|
||||
return []
|
||||
|
||||
async def clear_conversation(self, conversation_id: str) -> None:
|
||||
"""
|
||||
Clear all turns for a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
"""
|
||||
try:
|
||||
# Delete all points with this conversation_id
|
||||
self.client.delete(
|
||||
collection_name=self.collection_name,
|
||||
points_selector=Filter(
|
||||
must=[
|
||||
FieldCondition(
|
||||
key="conversation_id",
|
||||
match=MatchValue(value=conversation_id)
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
logger.info(f"Cleared conversation {conversation_id} from Qdrant")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error clearing conversation: {e}")
|
||||
raise
|
||||
|
||||
async def conversation_exists(self, conversation_id: str) -> bool:
|
||||
"""
|
||||
Check if a conversation exists
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
True if conversation has any turns
|
||||
"""
|
||||
try:
|
||||
points, _ = self.client.scroll(
|
||||
collection_name=self.collection_name,
|
||||
scroll_filter=Filter(
|
||||
must=[
|
||||
FieldCondition(
|
||||
key="conversation_id",
|
||||
match=MatchValue(value=conversation_id)
|
||||
)
|
||||
]
|
||||
),
|
||||
limit=1,
|
||||
with_payload=False,
|
||||
with_vectors=False
|
||||
)
|
||||
return len(points) > 0
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error checking conversation existence: {e}")
|
||||
return False
|
||||
|
||||
async def get_conversation_stats(self, conversation_id: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Get statistics about a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
Dictionary with stats
|
||||
"""
|
||||
try:
|
||||
points, _ = self.client.scroll(
|
||||
collection_name=self.collection_name,
|
||||
scroll_filter=Filter(
|
||||
must=[
|
||||
FieldCondition(
|
||||
key="conversation_id",
|
||||
match=MatchValue(value=conversation_id)
|
||||
)
|
||||
]
|
||||
),
|
||||
limit=1000, # Get all points
|
||||
with_payload=True,
|
||||
with_vectors=False
|
||||
)
|
||||
|
||||
total_turns = len(points)
|
||||
total_tokens = sum(
|
||||
point.payload.get("tokens_total", 0) for point in points
|
||||
)
|
||||
|
||||
return {
|
||||
"conversation_id": conversation_id,
|
||||
"total_turns": total_turns,
|
||||
"total_tokens": total_tokens,
|
||||
"exists": total_turns > 0
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting conversation stats: {e}")
|
||||
return {
|
||||
"conversation_id": conversation_id,
|
||||
"total_turns": 0,
|
||||
"total_tokens": 0,
|
||||
"exists": False
|
||||
}
|
||||
|
||||
|
||||
# Global instance
|
||||
_qdrant_memory: Optional[QdrantConversationMemory] = None
|
||||
|
||||
|
||||
def get_qdrant_memory() -> QdrantConversationMemory:
|
||||
"""
|
||||
Get or create global Qdrant memory instance
|
||||
|
||||
Returns:
|
||||
QdrantConversationMemory instance
|
||||
"""
|
||||
global _qdrant_memory
|
||||
if _qdrant_memory is None:
|
||||
_qdrant_memory = QdrantConversationMemory()
|
||||
return _qdrant_memory
|
||||
@@ -0,0 +1,109 @@
|
||||
"""
|
||||
Pydantic schemas for memory system
|
||||
"""
|
||||
from pydantic import BaseModel, Field
|
||||
from typing import List, Optional, Dict, Any
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class MessageRole(str, Enum):
|
||||
"""Message role types"""
|
||||
SYSTEM = "system"
|
||||
USER = "user"
|
||||
ASSISTANT = "assistant"
|
||||
|
||||
|
||||
class TokenUsage(BaseModel):
|
||||
"""Token usage information"""
|
||||
prompt: int = 0
|
||||
completion: int = 0
|
||||
total: int = 0
|
||||
|
||||
|
||||
class ConversationTurn(BaseModel):
|
||||
"""A single turn in a conversation"""
|
||||
role: MessageRole
|
||||
content: str
|
||||
timestamp: datetime = Field(default_factory=datetime.utcnow)
|
||||
turn_number: int
|
||||
tokens: Optional[TokenUsage] = None
|
||||
metadata: Dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class ConversationMetadata(BaseModel):
|
||||
"""Metadata about a conversation"""
|
||||
conversation_id: str
|
||||
user_id: Optional[str] = None
|
||||
created_at: datetime = Field(default_factory=datetime.utcnow)
|
||||
last_updated: datetime = Field(default_factory=datetime.utcnow)
|
||||
turn_count: int = 0
|
||||
total_tokens: int = 0
|
||||
status: str = "active" # active, archived, deleted
|
||||
|
||||
|
||||
class ConversationBuffer(BaseModel):
|
||||
"""In-memory conversation buffer (Tier 1)"""
|
||||
conversation_id: str
|
||||
turns: List[ConversationTurn] = Field(default_factory=list)
|
||||
metadata: ConversationMetadata
|
||||
|
||||
|
||||
class ConversationSummary(BaseModel):
|
||||
"""Summarized conversation segment (Tier 2)"""
|
||||
conversation_id: str
|
||||
summary_text: str
|
||||
turn_range_start: int
|
||||
turn_range_end: int
|
||||
created_at: datetime = Field(default_factory=datetime.utcnow)
|
||||
token_count: int = 0
|
||||
|
||||
|
||||
class MemoryQuery(BaseModel):
|
||||
"""Query for memory retrieval"""
|
||||
conversation_id: str
|
||||
query: Optional[str] = None
|
||||
limit: int = Field(default=10, ge=1, le=100)
|
||||
include_tier1: bool = True
|
||||
include_tier2: bool = True
|
||||
include_tier3: bool = True
|
||||
|
||||
|
||||
class MemoryResult(BaseModel):
|
||||
"""Result from memory retrieval"""
|
||||
conversation_id: str
|
||||
turns: List[ConversationTurn] = Field(default_factory=list)
|
||||
summaries: List[ConversationSummary] = Field(default_factory=list)
|
||||
source_tiers: List[int] = Field(default_factory=list) # Which tiers contributed
|
||||
total_results: int = 0
|
||||
|
||||
|
||||
# API Request/Response Models
|
||||
|
||||
class ConversationListResponse(BaseModel):
|
||||
"""Response for listing conversations"""
|
||||
conversations: List[ConversationMetadata]
|
||||
total: int
|
||||
page: int = 1
|
||||
page_size: int = 50
|
||||
|
||||
|
||||
class ConversationDetailResponse(BaseModel):
|
||||
"""Response for conversation details"""
|
||||
metadata: ConversationMetadata
|
||||
recent_turns: List[ConversationTurn]
|
||||
turn_count: int
|
||||
|
||||
|
||||
class ConversationSearchRequest(BaseModel):
|
||||
"""Request for semantic search in conversation"""
|
||||
query: str
|
||||
limit: int = Field(default=5, ge=1, le=50)
|
||||
|
||||
|
||||
class ConversationSearchResponse(BaseModel):
|
||||
"""Response for semantic search"""
|
||||
conversation_id: str
|
||||
results: List[ConversationTurn]
|
||||
scores: List[float] = Field(default_factory=list)
|
||||
total_results: int
|
||||
@@ -0,0 +1,239 @@
|
||||
"""
|
||||
Tier 1: ConversationBufferMemory (In-Memory Working Memory)
|
||||
|
||||
Fast in-memory storage for recent conversation turns.
|
||||
- Stores last N turns in RAM
|
||||
- < 1ms access time
|
||||
- Ephemeral (lost on restart)
|
||||
- Automatic pruning when limit reached
|
||||
"""
|
||||
import logging
|
||||
from typing import Dict, List, Optional
|
||||
from datetime import datetime
|
||||
from collections import OrderedDict
|
||||
|
||||
from .base import Tier1Memory
|
||||
from .schemas import (
|
||||
ConversationTurn,
|
||||
ConversationBuffer,
|
||||
ConversationMetadata,
|
||||
MessageRole,
|
||||
TokenUsage
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ConversationBufferMemory(Tier1Memory):
|
||||
"""
|
||||
In-memory buffer for recent conversation turns.
|
||||
|
||||
Stores the last N turns of each conversation in RAM for fast access.
|
||||
Automatically prunes old turns when limit is reached.
|
||||
"""
|
||||
|
||||
def __init__(self, max_turns: int = 10):
|
||||
"""
|
||||
Initialize buffer memory
|
||||
|
||||
Args:
|
||||
max_turns: Maximum number of turns to keep per conversation
|
||||
"""
|
||||
self.max_turns = max_turns
|
||||
# Use OrderedDict to maintain insertion order
|
||||
self._buffers: Dict[str, ConversationBuffer] = OrderedDict()
|
||||
logger.info(f"Initialized ConversationBufferMemory with max_turns={max_turns}")
|
||||
|
||||
async def add_turn(self, conversation_id: str, turn: ConversationTurn) -> None:
|
||||
"""
|
||||
Add a new turn to the buffer
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
turn: The conversation turn to store
|
||||
"""
|
||||
# Get or create buffer
|
||||
buffer = await self.get_buffer(conversation_id)
|
||||
if buffer is None:
|
||||
buffer = ConversationBuffer(
|
||||
conversation_id=conversation_id,
|
||||
turns=[],
|
||||
metadata=ConversationMetadata(
|
||||
conversation_id=conversation_id
|
||||
)
|
||||
)
|
||||
self._buffers[conversation_id] = buffer
|
||||
|
||||
# Add turn
|
||||
buffer.turns.append(turn)
|
||||
|
||||
# Update metadata
|
||||
buffer.metadata.turn_count = len(buffer.turns)
|
||||
buffer.metadata.last_updated = datetime.utcnow()
|
||||
|
||||
if turn.tokens:
|
||||
buffer.metadata.total_tokens += turn.tokens.total
|
||||
|
||||
# Auto-prune if exceeds max turns
|
||||
if len(buffer.turns) > self.max_turns:
|
||||
await self.prune(conversation_id, keep_last=self.max_turns)
|
||||
|
||||
logger.debug(
|
||||
f"Added turn {turn.turn_number} to conversation {conversation_id}. "
|
||||
f"Buffer size: {len(buffer.turns)}"
|
||||
)
|
||||
|
||||
async def get_turns(
|
||||
self,
|
||||
conversation_id: str,
|
||||
limit: Optional[int] = None,
|
||||
offset: int = 0
|
||||
) -> List[ConversationTurn]:
|
||||
"""
|
||||
Retrieve turns from the buffer
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
limit: Maximum number of turns to retrieve
|
||||
offset: Number of turns to skip
|
||||
|
||||
Returns:
|
||||
List of conversation turns
|
||||
"""
|
||||
buffer = await self.get_buffer(conversation_id)
|
||||
if buffer is None:
|
||||
return []
|
||||
|
||||
turns = buffer.turns[offset:]
|
||||
if limit:
|
||||
turns = turns[:limit]
|
||||
|
||||
return turns
|
||||
|
||||
async def get_recent_turns(
|
||||
self,
|
||||
conversation_id: str,
|
||||
limit: int = 10
|
||||
) -> List[ConversationTurn]:
|
||||
"""
|
||||
Get the most recent N turns
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
limit: Number of recent turns to retrieve
|
||||
|
||||
Returns:
|
||||
List of recent turns (most recent last)
|
||||
"""
|
||||
buffer = await self.get_buffer(conversation_id)
|
||||
if buffer is None:
|
||||
return []
|
||||
|
||||
return buffer.turns[-limit:] if len(buffer.turns) > limit else buffer.turns
|
||||
|
||||
async def get_buffer(self, conversation_id: str) -> Optional[ConversationBuffer]:
|
||||
"""
|
||||
Get the full conversation buffer
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
ConversationBuffer or None if not found
|
||||
"""
|
||||
return self._buffers.get(conversation_id)
|
||||
|
||||
async def clear_conversation(self, conversation_id: str) -> None:
|
||||
"""
|
||||
Clear all turns for a conversation
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
"""
|
||||
if conversation_id in self._buffers:
|
||||
del self._buffers[conversation_id]
|
||||
logger.info(f"Cleared buffer for conversation {conversation_id}")
|
||||
|
||||
async def conversation_exists(self, conversation_id: str) -> bool:
|
||||
"""
|
||||
Check if a conversation exists in the buffer
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
|
||||
Returns:
|
||||
True if conversation exists
|
||||
"""
|
||||
return conversation_id in self._buffers
|
||||
|
||||
async def prune(self, conversation_id: str, keep_last: int = 5) -> None:
|
||||
"""
|
||||
Prune old turns, keeping only the most recent ones
|
||||
|
||||
Args:
|
||||
conversation_id: Unique conversation identifier
|
||||
keep_last: Number of recent turns to keep
|
||||
"""
|
||||
buffer = await self.get_buffer(conversation_id)
|
||||
if buffer is None:
|
||||
return
|
||||
|
||||
if len(buffer.turns) > keep_last:
|
||||
removed_count = len(buffer.turns) - keep_last
|
||||
buffer.turns = buffer.turns[-keep_last:]
|
||||
buffer.metadata.turn_count = len(buffer.turns)
|
||||
|
||||
logger.debug(
|
||||
f"Pruned {removed_count} turns from conversation {conversation_id}. "
|
||||
f"Kept last {keep_last} turns."
|
||||
)
|
||||
|
||||
async def get_all_conversation_ids(self) -> List[str]:
|
||||
"""
|
||||
Get list of all conversation IDs in memory
|
||||
|
||||
Returns:
|
||||
List of conversation IDs
|
||||
"""
|
||||
return list(self._buffers.keys())
|
||||
|
||||
async def get_buffer_stats(self) -> dict:
|
||||
"""
|
||||
Get statistics about buffer memory usage
|
||||
|
||||
Returns:
|
||||
Dictionary with stats
|
||||
"""
|
||||
total_conversations = len(self._buffers)
|
||||
total_turns = sum(len(buf.turns) for buf in self._buffers.values())
|
||||
total_tokens = sum(buf.metadata.total_tokens for buf in self._buffers.values())
|
||||
|
||||
return {
|
||||
"total_conversations": total_conversations,
|
||||
"total_turns": total_turns,
|
||||
"total_tokens": total_tokens,
|
||||
"max_turns_per_conversation": self.max_turns,
|
||||
"avg_turns_per_conversation": (
|
||||
total_turns / total_conversations if total_conversations > 0 else 0
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
# Global instance
|
||||
_buffer_memory: Optional[ConversationBufferMemory] = None
|
||||
|
||||
|
||||
def get_buffer_memory(max_turns: int = 10) -> ConversationBufferMemory:
|
||||
"""
|
||||
Get or create the global buffer memory instance
|
||||
|
||||
Args:
|
||||
max_turns: Maximum turns per conversation
|
||||
|
||||
Returns:
|
||||
ConversationBufferMemory instance
|
||||
"""
|
||||
global _buffer_memory
|
||||
if _buffer_memory is None:
|
||||
_buffer_memory = ConversationBufferMemory(max_turns=max_turns)
|
||||
return _buffer_memory
|
||||
@@ -0,0 +1,128 @@
|
||||
"""
|
||||
Embedding model client for text vectorization
|
||||
|
||||
Uses sentence-transformers for generating embeddings.
|
||||
"""
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from src.config import get_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class EmbeddingClient:
|
||||
"""Client for generating text embeddings"""
|
||||
|
||||
def __init__(self, model_name: Optional[str] = None):
|
||||
"""
|
||||
Initialize embedding client
|
||||
|
||||
Args:
|
||||
model_name: Optional model name, defaults to config
|
||||
"""
|
||||
self.model_name = model_name or settings.embedding_model
|
||||
self.dimension = settings.embedding_dimension
|
||||
self._model: Optional[SentenceTransformer] = None
|
||||
logger.info(f"Initializing EmbeddingClient with model: {self.model_name}")
|
||||
|
||||
def _load_model(self) -> SentenceTransformer:
|
||||
"""
|
||||
Lazy load the embedding model
|
||||
|
||||
Returns:
|
||||
Loaded SentenceTransformer model
|
||||
"""
|
||||
if self._model is None:
|
||||
logger.info(f"Loading embedding model: {self.model_name}")
|
||||
self._model = SentenceTransformer(self.model_name)
|
||||
logger.info(f"Model loaded successfully. Embedding dimension: {self.dimension}")
|
||||
return self._model
|
||||
|
||||
def embed_text(self, text: str) -> List[float]:
|
||||
"""
|
||||
Generate embedding for a single text
|
||||
|
||||
Args:
|
||||
text: Input text to embed
|
||||
|
||||
Returns:
|
||||
List of floats representing the embedding vector
|
||||
"""
|
||||
model = self._load_model()
|
||||
embedding = model.encode(text, convert_to_numpy=True)
|
||||
return embedding.tolist()
|
||||
|
||||
def embed_batch(self, texts: List[str]) -> List[List[float]]:
|
||||
"""
|
||||
Generate embeddings for multiple texts
|
||||
|
||||
Args:
|
||||
texts: List of input texts
|
||||
|
||||
Returns:
|
||||
List of embedding vectors
|
||||
"""
|
||||
model = self._load_model()
|
||||
embeddings = model.encode(
|
||||
texts,
|
||||
batch_size=settings.embedding_batch_size,
|
||||
convert_to_numpy=True,
|
||||
show_progress_bar=False
|
||||
)
|
||||
return embeddings.tolist()
|
||||
|
||||
def get_dimension(self) -> int:
|
||||
"""
|
||||
Get embedding dimension
|
||||
|
||||
Returns:
|
||||
Embedding vector dimension
|
||||
"""
|
||||
return self.dimension
|
||||
|
||||
|
||||
# Global instance
|
||||
_embedding_client: Optional[EmbeddingClient] = None
|
||||
|
||||
|
||||
def get_embedding_client() -> EmbeddingClient:
|
||||
"""
|
||||
Get or create global embedding client instance
|
||||
|
||||
Returns:
|
||||
EmbeddingClient instance
|
||||
"""
|
||||
global _embedding_client
|
||||
if _embedding_client is None:
|
||||
_embedding_client = EmbeddingClient()
|
||||
return _embedding_client
|
||||
|
||||
|
||||
async def embed_text_async(text: str) -> List[float]:
|
||||
"""
|
||||
Async wrapper for embedding text
|
||||
|
||||
Args:
|
||||
text: Input text
|
||||
|
||||
Returns:
|
||||
Embedding vector
|
||||
"""
|
||||
client = get_embedding_client()
|
||||
return client.embed_text(text)
|
||||
|
||||
|
||||
async def embed_batch_async(texts: List[str]) -> List[List[float]]:
|
||||
"""
|
||||
Async wrapper for batch embedding
|
||||
|
||||
Args:
|
||||
texts: List of input texts
|
||||
|
||||
Returns:
|
||||
List of embedding vectors
|
||||
"""
|
||||
client = get_embedding_client()
|
||||
return client.embed_batch(texts)
|
||||
@@ -0,0 +1,201 @@
|
||||
"""
|
||||
Ollama client for model inference.
|
||||
Handles both streaming and non-streaming requests.
|
||||
"""
|
||||
|
||||
import httpx
|
||||
import json
|
||||
import logging
|
||||
from typing import AsyncIterator, Dict, Any, Optional
|
||||
from src.config import get_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
settings = get_settings()
|
||||
|
||||
|
||||
class OllamaClient:
|
||||
"""Client for interacting with Ollama API."""
|
||||
|
||||
def __init__(self):
|
||||
self.base_url = settings.ollama_base_url
|
||||
self.timeout = settings.ollama_timeout
|
||||
self.client = httpx.AsyncClient(timeout=self.timeout)
|
||||
logger.info(f"Initialized Ollama client: {self.base_url}")
|
||||
|
||||
async def close(self):
|
||||
"""Close the HTTP client."""
|
||||
await self.client.aclose()
|
||||
|
||||
def resolve_model(self, model_name: str) -> str:
|
||||
"""
|
||||
Resolve model alias to actual Ollama model.
|
||||
|
||||
Args:
|
||||
model_name: Requested model name (e.g., "gpt-3.5-turbo")
|
||||
|
||||
Returns:
|
||||
Actual Ollama model name (e.g., "gemma:7b")
|
||||
"""
|
||||
resolved = settings.model_aliases.get(model_name, model_name)
|
||||
if resolved != model_name:
|
||||
logger.info(f"Model resolution: {model_name} → {resolved}")
|
||||
return resolved
|
||||
|
||||
async def generate_non_streaming(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
temperature: float = 0.7,
|
||||
max_tokens: Optional[int] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Generate non-streaming response from Ollama.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
prompt: User prompt
|
||||
temperature: Sampling temperature
|
||||
max_tokens: Maximum tokens to generate
|
||||
|
||||
Returns:
|
||||
Dict with 'response' and 'tokens' keys
|
||||
"""
|
||||
actual_model = self.resolve_model(model)
|
||||
|
||||
payload = {
|
||||
"model": actual_model,
|
||||
"prompt": prompt,
|
||||
"stream": False,
|
||||
"options": {
|
||||
"temperature": temperature,
|
||||
}
|
||||
}
|
||||
|
||||
if max_tokens:
|
||||
payload["options"]["num_predict"] = max_tokens
|
||||
|
||||
logger.debug(f"Ollama request to {actual_model}")
|
||||
|
||||
try:
|
||||
response = await self.client.post(
|
||||
f"{self.base_url}/api/generate",
|
||||
json=payload
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
return {
|
||||
"response": result.get("response", ""),
|
||||
"tokens": {
|
||||
"prompt": result.get("prompt_eval_count", 0),
|
||||
"completion": result.get("eval_count", 0),
|
||||
"total": result.get("prompt_eval_count", 0) + result.get("eval_count", 0)
|
||||
}
|
||||
}
|
||||
|
||||
except httpx.HTTPError as e:
|
||||
logger.error(f"Ollama request failed: {e}")
|
||||
raise
|
||||
|
||||
async def generate_streaming(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
temperature: float = 0.7,
|
||||
max_tokens: Optional[int] = None
|
||||
) -> AsyncIterator[str]:
|
||||
"""
|
||||
Generate streaming response from Ollama.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
prompt: User prompt
|
||||
temperature: Sampling temperature
|
||||
max_tokens: Maximum tokens to generate
|
||||
|
||||
Yields:
|
||||
Token strings
|
||||
"""
|
||||
actual_model = self.resolve_model(model)
|
||||
|
||||
payload = {
|
||||
"model": actual_model,
|
||||
"prompt": prompt,
|
||||
"stream": True,
|
||||
"options": {
|
||||
"temperature": temperature,
|
||||
}
|
||||
}
|
||||
|
||||
if max_tokens:
|
||||
payload["options"]["num_predict"] = max_tokens
|
||||
|
||||
logger.debug(f"Ollama streaming request to {actual_model}")
|
||||
|
||||
try:
|
||||
async with self.client.stream(
|
||||
"POST",
|
||||
f"{self.base_url}/api/generate",
|
||||
json=payload
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
|
||||
async for line in response.aiter_lines():
|
||||
if not line:
|
||||
continue
|
||||
|
||||
try:
|
||||
chunk = json.loads(line)
|
||||
if "response" in chunk:
|
||||
token = chunk["response"]
|
||||
if token:
|
||||
yield token
|
||||
|
||||
# Check if done
|
||||
if chunk.get("done", False):
|
||||
break
|
||||
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(f"Failed to parse JSON: {line}")
|
||||
continue
|
||||
|
||||
except httpx.HTTPError as e:
|
||||
logger.error(f"Ollama streaming request failed: {e}")
|
||||
raise
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""
|
||||
Check if Ollama is healthy.
|
||||
|
||||
Returns:
|
||||
True if healthy, False otherwise
|
||||
"""
|
||||
try:
|
||||
response = await self.client.get(
|
||||
f"{self.base_url}/api/tags",
|
||||
timeout=5.0
|
||||
)
|
||||
return response.status_code == 200
|
||||
except Exception as e:
|
||||
logger.error(f"Ollama health check failed: {e}")
|
||||
return False
|
||||
|
||||
|
||||
# Global client instance
|
||||
_ollama_client: Optional[OllamaClient] = None
|
||||
|
||||
|
||||
def get_ollama_client() -> OllamaClient:
|
||||
"""Get or create the global Ollama client instance."""
|
||||
global _ollama_client
|
||||
if _ollama_client is None:
|
||||
_ollama_client = OllamaClient()
|
||||
return _ollama_client
|
||||
|
||||
|
||||
async def close_ollama_client():
|
||||
"""Close the global Ollama client."""
|
||||
global _ollama_client
|
||||
if _ollama_client is not None:
|
||||
await _ollama_client.close()
|
||||
_ollama_client = None
|
||||
@@ -0,0 +1,13 @@
|
||||
"""
|
||||
Web scraper module for extracting content from websites
|
||||
"""
|
||||
from src.web_scraper.router import router
|
||||
from src.web_scraper.schemas import WebScraperRequest, WebScraperResponse
|
||||
from src.web_scraper.service import WebScraperService
|
||||
|
||||
__all__ = [
|
||||
"router",
|
||||
"WebScraperRequest",
|
||||
"WebScraperResponse",
|
||||
"WebScraperService",
|
||||
]
|
||||
@@ -0,0 +1,32 @@
|
||||
"""
|
||||
Configuration for web scraper module
|
||||
"""
|
||||
from pydantic_settings import BaseSettings
|
||||
from functools import lru_cache
|
||||
|
||||
|
||||
class WebScraperSettings(BaseSettings):
|
||||
"""Web scraper specific settings"""
|
||||
|
||||
# HTTP client configuration
|
||||
request_timeout: int = 30
|
||||
max_redirects: int = 5
|
||||
user_agent: str = "Mozilla/5.0 (compatible; CoreCode/1.0)"
|
||||
|
||||
# Content extraction
|
||||
default_max_length: int = 10000
|
||||
max_links_to_extract: int = 50
|
||||
|
||||
# Rate limiting (future use)
|
||||
rate_limit_enabled: bool = False
|
||||
requests_per_minute: int = 60
|
||||
|
||||
class Config:
|
||||
env_prefix = "WEB_SCRAPER_"
|
||||
case_sensitive = False
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def get_web_scraper_settings() -> WebScraperSettings:
|
||||
"""Cached web scraper settings instance"""
|
||||
return WebScraperSettings()
|
||||
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
Custom exceptions for web scraper module
|
||||
"""
|
||||
|
||||
|
||||
class WebScraperException(Exception):
|
||||
"""Base exception for web scraper module"""
|
||||
pass
|
||||
|
||||
|
||||
class FetchError(WebScraperException):
|
||||
"""Raised when URL fetch fails"""
|
||||
pass
|
||||
|
||||
|
||||
class ScrapingError(WebScraperException):
|
||||
"""Raised when content extraction fails"""
|
||||
pass
|
||||
@@ -0,0 +1,79 @@
|
||||
"""
|
||||
API routes for web scraper module
|
||||
"""
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
|
||||
from src.logging_config import get_logger
|
||||
from src.web_scraper.schemas import WebScraperRequest, WebScraperResponse
|
||||
from src.web_scraper.service import WebScraperService
|
||||
from src.web_scraper.exceptions import FetchError, ScrapingError
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
router = APIRouter(
|
||||
prefix="/web-scraper",
|
||||
tags=["Web Scraper"]
|
||||
)
|
||||
|
||||
# Initialize service (could be dependency injected for testing)
|
||||
scraper_service = WebScraperService()
|
||||
|
||||
|
||||
@router.post(
|
||||
"/scrape",
|
||||
response_model=WebScraperResponse,
|
||||
status_code=status.HTTP_200_OK,
|
||||
summary="Scrape website content",
|
||||
description="""
|
||||
Scrape and extract main content from a website.
|
||||
|
||||
Uses trafilatura for intelligent content extraction (articles, blog posts, documentation),
|
||||
with BeautifulSoup as fallback. Perfect for feeding webpage content to LLMs.
|
||||
|
||||
**Features:**
|
||||
- Intelligent main content extraction
|
||||
- Removes navigation, ads, footers
|
||||
- Optional link extraction
|
||||
- Configurable content length limits
|
||||
|
||||
**Rate Limiting:** None (internal network use only)
|
||||
"""
|
||||
)
|
||||
async def scrape_website(request: WebScraperRequest) -> WebScraperResponse:
|
||||
"""
|
||||
Scrape a website and extract its main content
|
||||
|
||||
Args:
|
||||
request: Scraping request with URL and options
|
||||
|
||||
Returns:
|
||||
Extracted content with metadata
|
||||
|
||||
Raises:
|
||||
HTTPException: 400 for fetch errors, 500 for processing errors
|
||||
"""
|
||||
try:
|
||||
logger.info(f"Received scrape request for: {request.url}")
|
||||
result = await scraper_service.scrape_url(request)
|
||||
return result
|
||||
|
||||
except FetchError as e:
|
||||
logger.warning(f"Fetch failed: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Failed to fetch URL: {str(e)}"
|
||||
)
|
||||
|
||||
except ScrapingError as e:
|
||||
logger.error(f"Scraping failed: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to extract content: {str(e)}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Unexpected error: {str(e)}", exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="An unexpected error occurred"
|
||||
)
|
||||
@@ -0,0 +1,69 @@
|
||||
"""
|
||||
Pydantic schemas for web scraper module
|
||||
"""
|
||||
from pydantic import HttpUrl, Field
|
||||
from typing import Optional
|
||||
from datetime import datetime
|
||||
from src.base_schema import BaseSchema
|
||||
|
||||
|
||||
class WebScraperRequest(BaseSchema):
|
||||
"""Request model for web scraping"""
|
||||
|
||||
url: HttpUrl = Field(
|
||||
...,
|
||||
description="The URL to scrape",
|
||||
examples=["https://example.com/article"]
|
||||
)
|
||||
|
||||
extract_main_content: bool = Field(
|
||||
default=True,
|
||||
description="Use intelligent content extraction (trafilatura) vs raw HTML parsing"
|
||||
)
|
||||
|
||||
include_links: bool = Field(
|
||||
default=False,
|
||||
description="Include list of links found on the page"
|
||||
)
|
||||
|
||||
max_length: Optional[int] = Field(
|
||||
default=10000,
|
||||
ge=100,
|
||||
le=100000,
|
||||
description="Maximum content length to return (100-100000 chars)"
|
||||
)
|
||||
|
||||
|
||||
class WebScraperResponse(BaseSchema):
|
||||
"""Response model for web scraping"""
|
||||
|
||||
url: str = Field(
|
||||
...,
|
||||
description="The scraped URL"
|
||||
)
|
||||
|
||||
title: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Page title extracted from <title> tag"
|
||||
)
|
||||
|
||||
content: str = Field(
|
||||
...,
|
||||
description="Extracted page content"
|
||||
)
|
||||
|
||||
extracted_at: datetime = Field(
|
||||
...,
|
||||
description="UTC timestamp when content was extracted"
|
||||
)
|
||||
|
||||
content_length: int = Field(
|
||||
...,
|
||||
ge=0,
|
||||
description="Length of extracted content in characters"
|
||||
)
|
||||
|
||||
links: Optional[list[str]] = Field(
|
||||
default=None,
|
||||
description="List of HTTP(S) links found on the page (max 50)"
|
||||
)
|
||||
@@ -0,0 +1,213 @@
|
||||
"""
|
||||
Business logic for web scraper module
|
||||
"""
|
||||
import httpx
|
||||
from bs4 import BeautifulSoup
|
||||
import trafilatura
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from src.logging_config import get_logger
|
||||
from src.web_scraper.config import get_web_scraper_settings
|
||||
from src.web_scraper.schemas import WebScraperRequest, WebScraperResponse
|
||||
from src.web_scraper.exceptions import ScrapingError, FetchError
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class WebScraperService:
|
||||
"""Service class for web scraping operations"""
|
||||
|
||||
def __init__(self):
|
||||
self.settings = get_web_scraper_settings()
|
||||
|
||||
async def scrape_url(self, request: WebScraperRequest) -> WebScraperResponse:
|
||||
"""
|
||||
Scrape and extract content from a URL
|
||||
|
||||
Args:
|
||||
request: Scraping request parameters
|
||||
|
||||
Returns:
|
||||
Extracted content with metadata
|
||||
|
||||
Raises:
|
||||
FetchError: If URL cannot be fetched
|
||||
ScrapingError: If content extraction fails
|
||||
"""
|
||||
url_str = str(request.url)
|
||||
logger.info(f"Starting scrape for URL: {url_str}")
|
||||
|
||||
try:
|
||||
# Fetch the webpage
|
||||
html_content = await self._fetch_url(url_str)
|
||||
|
||||
# Extract content based on settings
|
||||
if request.extract_main_content:
|
||||
content = self._extract_main_content(html_content, request.include_links)
|
||||
else:
|
||||
content = self._extract_basic_content(html_content)
|
||||
|
||||
# Extract metadata
|
||||
title = self._extract_title(html_content)
|
||||
links = self._extract_links(html_content) if request.include_links else None
|
||||
|
||||
# Clean and truncate content
|
||||
content = self._clean_content(content)
|
||||
if request.max_length and len(content) > request.max_length:
|
||||
content = content[:request.max_length] + "\n\n[Content truncated...]"
|
||||
logger.debug(f"Content truncated to {request.max_length} characters")
|
||||
|
||||
logger.info(f"Successfully scraped {len(content)} characters from {url_str}")
|
||||
|
||||
return WebScraperResponse(
|
||||
url=url_str,
|
||||
title=title,
|
||||
content=content,
|
||||
extracted_at=datetime.now(timezone.utc),
|
||||
content_length=len(content),
|
||||
links=links
|
||||
)
|
||||
|
||||
except FetchError:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Scraping failed for {url_str}: {str(e)}", exc_info=True)
|
||||
raise ScrapingError(f"Failed to scrape content: {str(e)}")
|
||||
|
||||
async def _fetch_url(self, url: str) -> str:
|
||||
"""
|
||||
Fetch HTML content from URL
|
||||
|
||||
Args:
|
||||
url: URL to fetch
|
||||
|
||||
Returns:
|
||||
HTML content as string
|
||||
|
||||
Raises:
|
||||
FetchError: If fetch fails
|
||||
"""
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=self.settings.request_timeout,
|
||||
follow_redirects=True,
|
||||
max_redirects=self.settings.max_redirects
|
||||
) as client:
|
||||
logger.debug(f"Fetching URL: {url}")
|
||||
response = await client.get(
|
||||
url,
|
||||
headers={"User-Agent": self.settings.user_agent}
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.debug(f"Fetched {len(response.text)} bytes from {url}")
|
||||
return response.text
|
||||
|
||||
except httpx.HTTPStatusError as e:
|
||||
logger.error(f"HTTP error {e.response.status_code} for {url}")
|
||||
raise FetchError(f"HTTP {e.response.status_code}: {e.response.reason_phrase}")
|
||||
except httpx.RequestError as e:
|
||||
logger.error(f"Request error for {url}: {str(e)}")
|
||||
raise FetchError(f"Failed to fetch URL: {str(e)}")
|
||||
|
||||
def _extract_main_content(self, html: str, include_links: bool = False) -> str:
|
||||
"""
|
||||
Extract main content using trafilatura (intelligent extraction)
|
||||
|
||||
Args:
|
||||
html: Raw HTML content
|
||||
include_links: Whether to preserve links in output
|
||||
|
||||
Returns:
|
||||
Extracted content
|
||||
"""
|
||||
logger.debug("Extracting main content with trafilatura")
|
||||
content = trafilatura.extract(
|
||||
html,
|
||||
include_links=include_links,
|
||||
include_images=False,
|
||||
output_format='txt',
|
||||
no_fallback=False
|
||||
)
|
||||
|
||||
# Fallback to BeautifulSoup if trafilatura fails
|
||||
if not content:
|
||||
logger.debug("Trafilatura extraction failed, falling back to BeautifulSoup")
|
||||
content = self._extract_basic_content(html)
|
||||
|
||||
return content
|
||||
|
||||
def _extract_basic_content(self, html: str) -> str:
|
||||
"""
|
||||
Extract content using basic BeautifulSoup parsing
|
||||
|
||||
Args:
|
||||
html: Raw HTML content
|
||||
|
||||
Returns:
|
||||
Extracted text content
|
||||
"""
|
||||
logger.debug("Extracting content with BeautifulSoup")
|
||||
soup = BeautifulSoup(html, 'html.parser')
|
||||
|
||||
# Remove unwanted elements
|
||||
for element in soup(["script", "style", "nav", "footer", "header", "aside"]):
|
||||
element.decompose()
|
||||
|
||||
# Extract text
|
||||
text = soup.get_text(separator='\n', strip=True)
|
||||
return text
|
||||
|
||||
def _extract_title(self, html: str) -> Optional[str]:
|
||||
"""
|
||||
Extract page title from HTML
|
||||
|
||||
Args:
|
||||
html: Raw HTML content
|
||||
|
||||
Returns:
|
||||
Page title or None
|
||||
"""
|
||||
soup = BeautifulSoup(html, 'html.parser')
|
||||
title = soup.title.string if soup.title else None
|
||||
if title:
|
||||
title = title.strip()
|
||||
logger.debug(f"Extracted title: {title}")
|
||||
return title
|
||||
|
||||
def _extract_links(self, html: str) -> list[str]:
|
||||
"""
|
||||
Extract HTTP(S) links from HTML
|
||||
|
||||
Args:
|
||||
html: Raw HTML content
|
||||
|
||||
Returns:
|
||||
List of absolute HTTP(S) URLs
|
||||
"""
|
||||
soup = BeautifulSoup(html, 'html.parser')
|
||||
links = [
|
||||
a.get('href')
|
||||
for a in soup.find_all('a', href=True)
|
||||
if a.get('href', '').startswith('http')
|
||||
]
|
||||
|
||||
# Limit number of links
|
||||
links = links[:self.settings.max_links_to_extract]
|
||||
logger.debug(f"Extracted {len(links)} links")
|
||||
return links
|
||||
|
||||
def _clean_content(self, content: str) -> str:
|
||||
"""
|
||||
Clean and normalize extracted content
|
||||
|
||||
Args:
|
||||
content: Raw extracted content
|
||||
|
||||
Returns:
|
||||
Cleaned content
|
||||
"""
|
||||
# Remove empty lines and normalize whitespace
|
||||
lines = [line.strip() for line in content.split('\n') if line.strip()]
|
||||
cleaned = '\n'.join(lines)
|
||||
return cleaned
|
||||
@@ -0,0 +1,269 @@
|
||||
"""
|
||||
Integration tests for Phase 2 Memory System
|
||||
|
||||
Tests the complete memory stack:
|
||||
- Tier 1: ConversationBufferMemory
|
||||
- Tier 2/3: QdrantConversationMemory
|
||||
- Embedding Client
|
||||
"""
|
||||
import asyncio
|
||||
import pytest
|
||||
from datetime import datetime
|
||||
from src.memory import (
|
||||
ConversationBufferMemory,
|
||||
QdrantConversationMemory,
|
||||
ConversationTurn,
|
||||
MessageRole,
|
||||
TokenUsage,
|
||||
get_buffer_memory,
|
||||
get_qdrant_memory
|
||||
)
|
||||
from src.models.embeddings import get_embedding_client
|
||||
|
||||
|
||||
class TestEmbeddingClient:
|
||||
"""Test embedding generation"""
|
||||
|
||||
def test_embedding_client_init(self):
|
||||
"""Test embedding client initialization"""
|
||||
client = get_embedding_client()
|
||||
assert client is not None
|
||||
assert client.dimension == 384
|
||||
print(f"✓ Embedding client initialized: {client.model_name}")
|
||||
|
||||
def test_single_embedding(self):
|
||||
"""Test single text embedding"""
|
||||
client = get_embedding_client()
|
||||
text = "Hello, this is a test message for embedding generation"
|
||||
|
||||
embedding = client.embed_text(text)
|
||||
|
||||
assert isinstance(embedding, list)
|
||||
assert len(embedding) == 384
|
||||
assert all(isinstance(x, float) for x in embedding)
|
||||
print(f"✓ Single embedding generated: {len(embedding)} dimensions")
|
||||
|
||||
def test_batch_embedding(self):
|
||||
"""Test batch text embedding"""
|
||||
client = get_embedding_client()
|
||||
texts = [
|
||||
"First message about Python programming",
|
||||
"Second message about machine learning",
|
||||
"Third message about data science"
|
||||
]
|
||||
|
||||
embeddings = client.embed_batch(texts)
|
||||
|
||||
assert len(embeddings) == 3
|
||||
assert all(len(emb) == 384 for emb in embeddings)
|
||||
print(f"✓ Batch embeddings generated: {len(embeddings)} texts")
|
||||
|
||||
|
||||
class TestQdrantMemory:
|
||||
"""Test Qdrant memory storage and retrieval"""
|
||||
|
||||
@pytest.fixture
|
||||
def qdrant_memory(self):
|
||||
"""Get Qdrant memory instance"""
|
||||
return get_qdrant_memory()
|
||||
|
||||
@pytest.fixture
|
||||
def test_conversation_id(self):
|
||||
"""Generate unique test conversation ID"""
|
||||
return f"test_conv_{int(datetime.utcnow().timestamp())}"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_qdrant_connection(self, qdrant_memory):
|
||||
"""Test Qdrant connection and collection"""
|
||||
assert qdrant_memory.client is not None
|
||||
assert qdrant_memory.collection_name == "core_api_conversations"
|
||||
print(f"✓ Connected to Qdrant: {qdrant_memory.host}:{qdrant_memory.port}")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_turn(self, qdrant_memory, test_conversation_id):
|
||||
"""Test adding a turn to Qdrant"""
|
||||
turn = ConversationTurn(
|
||||
role=MessageRole.USER,
|
||||
content="What is Python?",
|
||||
turn_number=1,
|
||||
tokens=TokenUsage(prompt=10, completion=0, total=10)
|
||||
)
|
||||
|
||||
await qdrant_memory.add_turn(test_conversation_id, turn)
|
||||
|
||||
# Verify it was stored
|
||||
exists = await qdrant_memory.conversation_exists(test_conversation_id)
|
||||
assert exists is True
|
||||
print(f"✓ Turn stored in Qdrant: {test_conversation_id}")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chronological_retrieval(self, qdrant_memory, test_conversation_id):
|
||||
"""Test Tier 2 mode: chronological retrieval"""
|
||||
# Add multiple turns
|
||||
turns = [
|
||||
ConversationTurn(role=MessageRole.USER, content="What is Python?", turn_number=1),
|
||||
ConversationTurn(role=MessageRole.ASSISTANT, content="Python is a programming language", turn_number=2),
|
||||
ConversationTurn(role=MessageRole.USER, content="How do I learn it?", turn_number=3),
|
||||
]
|
||||
|
||||
for turn in turns:
|
||||
await qdrant_memory.add_turn(test_conversation_id, turn)
|
||||
|
||||
# Retrieve turns chronologically
|
||||
retrieved = await qdrant_memory.get_turns(test_conversation_id)
|
||||
|
||||
assert len(retrieved) == 3
|
||||
assert retrieved[0].turn_number == 1
|
||||
assert retrieved[1].turn_number == 2
|
||||
assert retrieved[2].turn_number == 3
|
||||
assert retrieved[0].content == "What is Python?"
|
||||
print(f"✓ Chronological retrieval works: {len(retrieved)} turns")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_semantic_search(self, qdrant_memory, test_conversation_id):
|
||||
"""Test Tier 3 mode: semantic search"""
|
||||
# Add turns with distinct topics
|
||||
turns = [
|
||||
ConversationTurn(role=MessageRole.USER, content="I love machine learning and neural networks", turn_number=10),
|
||||
ConversationTurn(role=MessageRole.USER, content="Pizza is my favorite food", turn_number=11),
|
||||
ConversationTurn(role=MessageRole.USER, content="Deep learning models are fascinating", turn_number=12),
|
||||
]
|
||||
|
||||
for turn in turns:
|
||||
await qdrant_memory.add_turn(test_conversation_id, turn)
|
||||
|
||||
# Search for AI-related content
|
||||
results = await qdrant_memory.similarity_search(
|
||||
query="artificial intelligence and AI",
|
||||
conversation_id=test_conversation_id,
|
||||
limit=3
|
||||
)
|
||||
|
||||
assert len(results) > 0
|
||||
# Top results should be about ML/AI, not pizza
|
||||
top_result = results[0]
|
||||
assert "machine learning" in top_result["content"] or "Deep learning" in top_result["content"]
|
||||
assert top_result["score"] > 0.5 # Reasonable similarity score
|
||||
print(f"✓ Semantic search works: {len(results)} matches, top score: {results[0]['score']:.3f}")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_conversation_stats(self, qdrant_memory, test_conversation_id):
|
||||
"""Test conversation statistics"""
|
||||
stats = await qdrant_memory.get_conversation_stats(test_conversation_id)
|
||||
|
||||
assert stats["conversation_id"] == test_conversation_id
|
||||
assert stats["total_turns"] >= 0
|
||||
assert "total_tokens" in stats
|
||||
print(f"✓ Stats retrieved: {stats['total_turns']} turns, {stats['total_tokens']} tokens")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clear_conversation(self, qdrant_memory, test_conversation_id):
|
||||
"""Test clearing a conversation"""
|
||||
# Add a turn
|
||||
turn = ConversationTurn(role=MessageRole.USER, content="Test message", turn_number=99)
|
||||
await qdrant_memory.add_turn(test_conversation_id, turn)
|
||||
|
||||
# Clear it
|
||||
await qdrant_memory.clear_conversation(test_conversation_id)
|
||||
|
||||
# Verify it's gone
|
||||
exists = await qdrant_memory.conversation_exists(test_conversation_id)
|
||||
assert exists is False
|
||||
print(f"✓ Conversation cleared: {test_conversation_id}")
|
||||
|
||||
|
||||
class TestIntegration:
|
||||
"""Test full integration: Tier 1 + Qdrant + Embeddings"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_memory_flow(self):
|
||||
"""Test complete memory flow: Buffer → Qdrant"""
|
||||
conversation_id = f"integration_test_{int(datetime.utcnow().timestamp())}"
|
||||
|
||||
# Initialize both tiers
|
||||
buffer_memory = get_buffer_memory()
|
||||
qdrant_memory = get_qdrant_memory()
|
||||
|
||||
# 1. Add turns to buffer (Tier 1)
|
||||
turns = [
|
||||
ConversationTurn(role=MessageRole.USER, content="Hello!", turn_number=1),
|
||||
ConversationTurn(role=MessageRole.ASSISTANT, content="Hi there!", turn_number=2),
|
||||
ConversationTurn(role=MessageRole.USER, content="How are you?", turn_number=3),
|
||||
]
|
||||
|
||||
for turn in turns:
|
||||
await buffer_memory.add_turn(conversation_id, turn)
|
||||
|
||||
# Verify buffer has them
|
||||
buffer = await buffer_memory.get_buffer(conversation_id)
|
||||
assert len(buffer.turns) == 3
|
||||
print(f"✓ Tier 1 buffer: {len(buffer.turns)} turns")
|
||||
|
||||
# 2. Move to Qdrant (Tier 2/3)
|
||||
for turn in buffer.turns:
|
||||
await qdrant_memory.add_turn(conversation_id, turn)
|
||||
|
||||
# Verify Qdrant has them
|
||||
qdrant_turns = await qdrant_memory.get_turns(conversation_id)
|
||||
assert len(qdrant_turns) == 3
|
||||
print(f"✓ Tier 2/3 Qdrant: {len(qdrant_turns)} turns")
|
||||
|
||||
# 3. Test semantic search across both
|
||||
search_results = await qdrant_memory.similarity_search(
|
||||
query="greeting",
|
||||
conversation_id=conversation_id,
|
||||
limit=2
|
||||
)
|
||||
assert len(search_results) > 0
|
||||
print(f"✓ Semantic search: {len(search_results)} matches")
|
||||
|
||||
# Cleanup
|
||||
await qdrant_memory.clear_conversation(conversation_id)
|
||||
await buffer_memory.clear_conversation(conversation_id)
|
||||
print(f"✓ Full memory flow complete!")
|
||||
|
||||
|
||||
def run_tests():
|
||||
"""Run all tests"""
|
||||
print("\n" + "="*60)
|
||||
print("Phase 2 Memory System Integration Tests")
|
||||
print("="*60 + "\n")
|
||||
|
||||
# Test 1: Embedding Client
|
||||
print("Test 1: Embedding Client")
|
||||
print("-" * 40)
|
||||
test_embed = TestEmbeddingClient()
|
||||
test_embed.test_embedding_client_init()
|
||||
test_embed.test_single_embedding()
|
||||
test_embed.test_batch_embedding()
|
||||
print()
|
||||
|
||||
# Test 2: Qdrant Memory
|
||||
print("Test 2: Qdrant Memory Storage")
|
||||
print("-" * 40)
|
||||
test_qdrant = TestQdrantMemory()
|
||||
qdrant_memory = get_qdrant_memory()
|
||||
test_conv_id = f"test_conv_{int(datetime.utcnow().timestamp())}"
|
||||
|
||||
asyncio.run(test_qdrant.test_qdrant_connection(qdrant_memory))
|
||||
asyncio.run(test_qdrant.test_add_turn(qdrant_memory, test_conv_id))
|
||||
asyncio.run(test_qdrant.test_chronological_retrieval(qdrant_memory, test_conv_id))
|
||||
asyncio.run(test_qdrant.test_semantic_search(qdrant_memory, test_conv_id))
|
||||
asyncio.run(test_qdrant.test_conversation_stats(qdrant_memory, test_conv_id))
|
||||
asyncio.run(test_qdrant.test_clear_conversation(qdrant_memory, test_conv_id))
|
||||
print()
|
||||
|
||||
# Test 3: Full Integration
|
||||
print("Test 3: Full Integration (Tier 1 + Tier 2/3)")
|
||||
print("-" * 40)
|
||||
test_integration = TestIntegration()
|
||||
asyncio.run(test_integration.test_full_memory_flow())
|
||||
print()
|
||||
|
||||
print("="*60)
|
||||
print("✅ All Memory System Tests Passed!")
|
||||
print("="*60)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run_tests()
|
||||
@@ -0,0 +1,150 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Test MemoryManager orchestration
|
||||
|
||||
Verifies unified memory interface works correctly.
|
||||
"""
|
||||
import asyncio
|
||||
import sys
|
||||
from datetime import datetime
|
||||
|
||||
sys.path.insert(0, '/app')
|
||||
|
||||
from src.memory import MemoryManager, get_memory_manager, MessageRole, TokenUsage
|
||||
|
||||
|
||||
async def test_memory_manager():
|
||||
"""Test MemoryManager orchestration"""
|
||||
print("\n" + "="*60)
|
||||
print("MEMORY MANAGER TEST")
|
||||
print("="*60)
|
||||
|
||||
test_conv_id = f"manager_test_{int(datetime.utcnow().timestamp())}"
|
||||
|
||||
try:
|
||||
# Initialize manager
|
||||
manager = get_memory_manager()
|
||||
print(f"✓ MemoryManager initialized")
|
||||
|
||||
# Test 1: Add turns through manager
|
||||
print("\n1. Adding turns via MemoryManager...")
|
||||
turn1 = await manager.add_turn(
|
||||
conversation_id=test_conv_id,
|
||||
role=MessageRole.USER,
|
||||
content="Hello, how are you?",
|
||||
tokens=TokenUsage(prompt=5, completion=0, total=5)
|
||||
)
|
||||
assert turn1.turn_number == 1
|
||||
print(f" ✓ Turn 1 added: {turn1.content[:30]}...")
|
||||
|
||||
turn2 = await manager.add_turn(
|
||||
conversation_id=test_conv_id,
|
||||
role=MessageRole.ASSISTANT,
|
||||
content="I'm doing great! How can I help you today?",
|
||||
tokens=TokenUsage(prompt=5, completion=10, total=15)
|
||||
)
|
||||
assert turn2.turn_number == 2
|
||||
print(f" ✓ Turn 2 added: {turn2.content[:30]}...")
|
||||
|
||||
# Test 2: Get recent turns (from buffer)
|
||||
print("\n2. Getting recent turns from buffer...")
|
||||
recent = await manager.get_recent_turns(test_conv_id, limit=10)
|
||||
assert len(recent) == 2
|
||||
assert recent[0].turn_number == 1
|
||||
assert recent[1].turn_number == 2
|
||||
print(f" ✓ Retrieved {len(recent)} recent turns from buffer")
|
||||
|
||||
# Test 3: Add more turns to trigger consolidation (threshold = 10)
|
||||
print("\n3. Adding turns to trigger auto-consolidation...")
|
||||
for i in range(3, 11): # Add turns 3-10
|
||||
await manager.add_turn(
|
||||
conversation_id=test_conv_id,
|
||||
role=MessageRole.USER if i % 2 == 1 else MessageRole.ASSISTANT,
|
||||
content=f"Test message number {i}",
|
||||
tokens=TokenUsage(prompt=5, completion=5, total=10)
|
||||
)
|
||||
print(f" ✓ Added 8 more turns (total: 10)")
|
||||
|
||||
# Check if consolidation happened (turn 10 should trigger it)
|
||||
print("\n4. Verifying auto-consolidation...")
|
||||
stats = await manager.get_conversation_stats(test_conv_id)
|
||||
print(f" Buffer turns: {stats['buffer_turns']}")
|
||||
print(f" Qdrant turns: {stats['qdrant_turns']}")
|
||||
print(f" Exists in buffer: {stats['exists_in_buffer']}")
|
||||
print(f" Exists in Qdrant: {stats['exists_in_qdrant']}")
|
||||
|
||||
if stats['qdrant_turns'] > 0:
|
||||
print(f" ✓ Auto-consolidation triggered! {stats['qdrant_turns']} turns in Qdrant")
|
||||
else:
|
||||
print(f" ⚠ No auto-consolidation yet (threshold may not be reached)")
|
||||
|
||||
# Test 4: Manual consolidation
|
||||
print("\n5. Testing manual consolidation...")
|
||||
consolidated = await manager.consolidate(test_conv_id)
|
||||
print(f" ✓ Manually consolidated {consolidated} turns")
|
||||
|
||||
# Test 5: Get full history (buffer + Qdrant)
|
||||
print("\n6. Getting full conversation history...")
|
||||
full_history = await manager.get_full_history(test_conv_id)
|
||||
print(f" ✓ Retrieved {len(full_history)} total turns")
|
||||
assert len(full_history) == 10, f"Expected 10 turns, got {len(full_history)}"
|
||||
print(f" ✓ Full history verified (10 turns)")
|
||||
|
||||
# Test 6: Semantic search
|
||||
print("\n7. Testing semantic search...")
|
||||
search_results = await manager.search_conversations(
|
||||
query="greeting hello",
|
||||
conversation_id=test_conv_id,
|
||||
limit=3
|
||||
)
|
||||
if len(search_results) > 0:
|
||||
print(f" ✓ Semantic search found {len(search_results)} matches")
|
||||
print(f" Top: '{search_results[0]['content'][:40]}...' (score: {search_results[0]['score']:.3f})")
|
||||
else:
|
||||
print(f" ⚠ No semantic search results (may need more data)")
|
||||
|
||||
# Test 7: Clear conversation
|
||||
print("\n8. Clearing conversation...")
|
||||
await manager.clear_conversation(test_conv_id)
|
||||
stats_after = await manager.get_conversation_stats(test_conv_id)
|
||||
assert stats_after['buffer_turns'] == 0
|
||||
assert stats_after['qdrant_turns'] == 0
|
||||
print(f" ✓ Conversation cleared from all tiers")
|
||||
|
||||
print("\n" + "="*60)
|
||||
print("✅ MEMORY MANAGER TEST: PASSED")
|
||||
print("="*60)
|
||||
print("\nMemoryManager verified:")
|
||||
print(" ✓ Add turns with auto turn numbering")
|
||||
print(" ✓ Get recent turns from buffer")
|
||||
print(" ✓ Auto-consolidation (when threshold reached)")
|
||||
print(" ✓ Manual consolidation")
|
||||
print(" ✓ Get full history (buffer + Qdrant)")
|
||||
print(" ✓ Semantic search")
|
||||
print(" ✓ Clear conversation")
|
||||
print(" ✓ Conversation stats")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(f"\n❌ MEMORY MANAGER TEST: FAILED")
|
||||
print(f"Error: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
# Cleanup on error
|
||||
try:
|
||||
await manager.clear_conversation(test_conv_id)
|
||||
except:
|
||||
pass
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def main():
|
||||
"""Run the test"""
|
||||
success = asyncio.run(test_memory_manager())
|
||||
return 0 if success else 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
exit(main())
|
||||
@@ -0,0 +1,330 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Simple integration tests for Phase 2 Memory System
|
||||
No external dependencies beyond the memory system itself
|
||||
"""
|
||||
import asyncio
|
||||
import sys
|
||||
from datetime import datetime
|
||||
|
||||
# Add src to path
|
||||
sys.path.insert(0, '/app')
|
||||
|
||||
from src.memory import (
|
||||
ConversationBufferMemory,
|
||||
QdrantConversationMemory,
|
||||
ConversationTurn,
|
||||
MessageRole,
|
||||
TokenUsage,
|
||||
get_buffer_memory,
|
||||
get_qdrant_memory
|
||||
)
|
||||
from src.models.embeddings import get_embedding_client
|
||||
|
||||
|
||||
def test_embedding_client():
|
||||
"""Test 1: Embedding Client"""
|
||||
print("\n" + "="*60)
|
||||
print("Test 1: Embedding Client")
|
||||
print("="*60)
|
||||
|
||||
try:
|
||||
# Initialize
|
||||
client = get_embedding_client()
|
||||
assert client is not None
|
||||
assert client.dimension == 384
|
||||
print(f"✓ Embedding client initialized: {client.model_name}")
|
||||
print(f"✓ Embedding dimension: {client.dimension}")
|
||||
|
||||
# Single embedding
|
||||
text = "Hello, this is a test message for embedding generation"
|
||||
embedding = client.embed_text(text)
|
||||
assert isinstance(embedding, list)
|
||||
assert len(embedding) == 384
|
||||
assert all(isinstance(x, float) for x in embedding)
|
||||
print(f"✓ Single embedding generated: {len(embedding)} dimensions")
|
||||
print(f" Sample values: [{embedding[0]:.4f}, {embedding[1]:.4f}, {embedding[2]:.4f}, ...]")
|
||||
|
||||
# Batch embedding
|
||||
texts = [
|
||||
"First message about Python programming",
|
||||
"Second message about machine learning",
|
||||
"Third message about data science"
|
||||
]
|
||||
embeddings = client.embed_batch(texts)
|
||||
assert len(embeddings) == 3
|
||||
assert all(len(emb) == 384 for emb in embeddings)
|
||||
print(f"✓ Batch embeddings generated: {len(embeddings)} texts")
|
||||
|
||||
print("\n✅ Embedding Client Tests: PASSED")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(f"\n❌ Embedding Client Tests: FAILED")
|
||||
print(f"Error: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return False
|
||||
|
||||
|
||||
async def test_qdrant_memory():
|
||||
"""Test 2: Qdrant Memory Storage"""
|
||||
print("\n" + "="*60)
|
||||
print("Test 2: Qdrant Memory Storage")
|
||||
print("="*60)
|
||||
|
||||
test_conv_id = f"test_conv_{int(datetime.utcnow().timestamp())}"
|
||||
|
||||
try:
|
||||
# Initialize
|
||||
qdrant_memory = get_qdrant_memory()
|
||||
assert qdrant_memory.client is not None
|
||||
assert qdrant_memory.collection_name == "core_api_conversations"
|
||||
print(f"✓ Connected to Qdrant: {qdrant_memory.host}:{qdrant_memory.port}")
|
||||
print(f"✓ Collection: {qdrant_memory.collection_name}")
|
||||
|
||||
# Add single turn
|
||||
turn1 = ConversationTurn(
|
||||
role=MessageRole.USER,
|
||||
content="What is Python?",
|
||||
turn_number=1,
|
||||
tokens=TokenUsage(prompt=10, completion=0, total=10)
|
||||
)
|
||||
await qdrant_memory.add_turn(test_conv_id, turn1)
|
||||
print(f"✓ Turn 1 stored in Qdrant")
|
||||
|
||||
# Verify it exists
|
||||
exists = await qdrant_memory.conversation_exists(test_conv_id)
|
||||
assert exists is True
|
||||
print(f"✓ Conversation exists: {test_conv_id}")
|
||||
|
||||
# Add more turns for chronological test
|
||||
turn2 = ConversationTurn(
|
||||
role=MessageRole.ASSISTANT,
|
||||
content="Python is a high-level programming language known for simplicity and readability",
|
||||
turn_number=2
|
||||
)
|
||||
turn3 = ConversationTurn(
|
||||
role=MessageRole.USER,
|
||||
content="How do I learn Python programming?",
|
||||
turn_number=3
|
||||
)
|
||||
|
||||
await qdrant_memory.add_turn(test_conv_id, turn2)
|
||||
await qdrant_memory.add_turn(test_conv_id, turn3)
|
||||
print(f"✓ Turns 2-3 stored in Qdrant")
|
||||
|
||||
# Test chronological retrieval (Tier 2 mode)
|
||||
retrieved = await qdrant_memory.get_turns(test_conv_id)
|
||||
assert len(retrieved) == 3
|
||||
assert retrieved[0].turn_number == 1
|
||||
assert retrieved[1].turn_number == 2
|
||||
assert retrieved[2].turn_number == 3
|
||||
assert retrieved[0].content == "What is Python?"
|
||||
print(f"✓ Chronological retrieval works: {len(retrieved)} turns")
|
||||
for i, turn in enumerate(retrieved, 1):
|
||||
print(f" Turn {turn.turn_number}: {turn.role.value} - {turn.content[:50]}...")
|
||||
|
||||
# Add turns with distinct topics for semantic search
|
||||
turn10 = ConversationTurn(
|
||||
role=MessageRole.USER,
|
||||
content="I love machine learning and neural networks and artificial intelligence",
|
||||
turn_number=10
|
||||
)
|
||||
turn11 = ConversationTurn(
|
||||
role=MessageRole.USER,
|
||||
content="Pizza is my favorite food and I enjoy eating pasta",
|
||||
turn_number=11
|
||||
)
|
||||
turn12 = ConversationTurn(
|
||||
role=MessageRole.USER,
|
||||
content="Deep learning models and transformers are fascinating AI technologies",
|
||||
turn_number=12
|
||||
)
|
||||
|
||||
await qdrant_memory.add_turn(test_conv_id, turn10)
|
||||
await qdrant_memory.add_turn(test_conv_id, turn11)
|
||||
await qdrant_memory.add_turn(test_conv_id, turn12)
|
||||
print(f"✓ Added 3 more turns for semantic search test")
|
||||
|
||||
# Test semantic search (Tier 3 mode)
|
||||
search_results = await qdrant_memory.similarity_search(
|
||||
query="artificial intelligence and deep learning",
|
||||
conversation_id=test_conv_id,
|
||||
limit=3
|
||||
)
|
||||
assert len(search_results) > 0
|
||||
print(f"✓ Semantic search works: {len(search_results)} matches")
|
||||
|
||||
# Top result should be about AI/ML, not food
|
||||
top_result = search_results[0]
|
||||
print(f" Top match (score: {top_result['score']:.3f}): {top_result['content'][:60]}...")
|
||||
assert top_result["score"] > 0.5, "Semantic similarity score too low"
|
||||
|
||||
# Verify top matches are AI-related
|
||||
ai_keywords = ["machine learning", "neural networks", "Deep learning", "AI", "artificial intelligence"]
|
||||
top_content = search_results[0]["content"]
|
||||
assert any(keyword in top_content for keyword in ai_keywords), "Top result not AI-related"
|
||||
print(f"✓ Semantic relevance verified (AI-related content ranked higher)")
|
||||
|
||||
# Test conversation stats
|
||||
stats = await qdrant_memory.get_conversation_stats(test_conv_id)
|
||||
assert stats["conversation_id"] == test_conv_id
|
||||
assert stats["total_turns"] == 6
|
||||
print(f"✓ Stats retrieved: {stats['total_turns']} turns, {stats['total_tokens']} tokens")
|
||||
|
||||
# Cleanup
|
||||
await qdrant_memory.clear_conversation(test_conv_id)
|
||||
exists_after = await qdrant_memory.conversation_exists(test_conv_id)
|
||||
assert exists_after is False
|
||||
print(f"✓ Conversation cleared successfully")
|
||||
|
||||
print("\n✅ Qdrant Memory Tests: PASSED")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(f"\n❌ Qdrant Memory Tests: FAILED")
|
||||
print(f"Error: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
# Cleanup on error
|
||||
try:
|
||||
await qdrant_memory.clear_conversation(test_conv_id)
|
||||
except:
|
||||
pass
|
||||
|
||||
return False
|
||||
|
||||
|
||||
async def test_full_integration():
|
||||
"""Test 3: Full Integration (Tier 1 + Tier 2/3)"""
|
||||
print("\n" + "="*60)
|
||||
print("Test 3: Full Integration (Tier 1 + Tier 2/3)")
|
||||
print("="*60)
|
||||
|
||||
test_conv_id = f"integration_test_{int(datetime.utcnow().timestamp())}"
|
||||
|
||||
try:
|
||||
# Initialize both tiers
|
||||
buffer_memory = get_buffer_memory()
|
||||
qdrant_memory = get_qdrant_memory()
|
||||
print(f"✓ Initialized Tier 1 (Buffer) and Tier 2/3 (Qdrant)")
|
||||
|
||||
# 1. Add turns to buffer (Tier 1)
|
||||
turns = [
|
||||
ConversationTurn(role=MessageRole.USER, content="Hello!", turn_number=1),
|
||||
ConversationTurn(role=MessageRole.ASSISTANT, content="Hi there! How can I help?", turn_number=2),
|
||||
ConversationTurn(role=MessageRole.USER, content="How are you?", turn_number=3),
|
||||
ConversationTurn(role=MessageRole.ASSISTANT, content="I'm doing great, thanks!", turn_number=4),
|
||||
]
|
||||
|
||||
for turn in turns:
|
||||
await buffer_memory.add_turn(test_conv_id, turn)
|
||||
|
||||
# Verify buffer has them
|
||||
buffer = await buffer_memory.get_buffer(test_conv_id)
|
||||
assert len(buffer.turns) == 4
|
||||
print(f"✓ Tier 1 buffer: {len(buffer.turns)} turns stored")
|
||||
|
||||
# 2. Move to Qdrant (Tier 2/3) - simulating consolidation
|
||||
for turn in buffer.turns:
|
||||
await qdrant_memory.add_turn(test_conv_id, turn)
|
||||
|
||||
# Verify Qdrant has them
|
||||
qdrant_turns = await qdrant_memory.get_turns(test_conv_id)
|
||||
assert len(qdrant_turns) == 4
|
||||
print(f"✓ Tier 2/3 Qdrant: {len(qdrant_turns)} turns stored")
|
||||
|
||||
# 3. Test semantic search across consolidated data
|
||||
search_results = await qdrant_memory.similarity_search(
|
||||
query="greeting hello",
|
||||
conversation_id=test_conv_id,
|
||||
limit=2
|
||||
)
|
||||
assert len(search_results) > 0
|
||||
print(f"✓ Semantic search: {len(search_results)} matches found")
|
||||
print(f" Best match: '{search_results[0]['content']}' (score: {search_results[0]['score']:.3f})")
|
||||
|
||||
# 4. Test data consistency
|
||||
buffer_content = [t.content for t in buffer.turns]
|
||||
qdrant_content = [t.content for t in qdrant_turns]
|
||||
assert buffer_content == qdrant_content
|
||||
print(f"✓ Data consistency verified (Buffer ↔ Qdrant)")
|
||||
|
||||
# Cleanup
|
||||
await qdrant_memory.clear_conversation(test_conv_id)
|
||||
await buffer_memory.clear_conversation(test_conv_id)
|
||||
print(f"✓ Cleanup complete")
|
||||
|
||||
print("\n✅ Full Integration Tests: PASSED")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(f"\n❌ Full Integration Tests: FAILED")
|
||||
print(f"Error: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
|
||||
# Cleanup on error
|
||||
try:
|
||||
await qdrant_memory.clear_conversation(test_conv_id)
|
||||
await buffer_memory.clear_conversation(test_conv_id)
|
||||
except:
|
||||
pass
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def main():
|
||||
"""Run all tests"""
|
||||
print("\n" + "="*60)
|
||||
print("PHASE 2 MEMORY SYSTEM - INTEGRATION TESTS")
|
||||
print("="*60)
|
||||
print(f"Start time: {datetime.utcnow().isoformat()}")
|
||||
|
||||
results = []
|
||||
|
||||
# Test 1: Embedding Client
|
||||
results.append(("Embedding Client", test_embedding_client()))
|
||||
|
||||
# Test 2: Qdrant Memory
|
||||
results.append(("Qdrant Memory", asyncio.run(test_qdrant_memory())))
|
||||
|
||||
# Test 3: Full Integration
|
||||
results.append(("Full Integration", asyncio.run(test_full_integration())))
|
||||
|
||||
# Summary
|
||||
print("\n" + "="*60)
|
||||
print("TEST SUMMARY")
|
||||
print("="*60)
|
||||
|
||||
for test_name, passed in results:
|
||||
status = "✅ PASSED" if passed else "❌ FAILED"
|
||||
print(f"{test_name:.<40} {status}")
|
||||
|
||||
total = len(results)
|
||||
passed = sum(1 for _, p in results if p)
|
||||
failed = total - passed
|
||||
|
||||
print(f"\nTotal: {total} | Passed: {passed} | Failed: {failed}")
|
||||
print(f"Success rate: {(passed/total)*100:.1f}%")
|
||||
|
||||
if all(p for _, p in results):
|
||||
print("\n" + "="*60)
|
||||
print("🎉 ALL TESTS PASSED!")
|
||||
print("="*60)
|
||||
print("\nPhase 2 Memory System Status: ✅ FUNCTIONAL")
|
||||
print("- Embedding client working (384d vectors)")
|
||||
print("- Qdrant storage working (chronological + semantic)")
|
||||
print("- Full integration working (Tier 1 ↔ Tier 2/3)")
|
||||
return 0
|
||||
else:
|
||||
print("\n" + "="*60)
|
||||
print("❌ SOME TESTS FAILED")
|
||||
print("="*60)
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
exit(main())
|
||||
Reference in New Issue
Block a user