Spaces:
Running
Running
add aplication file
Browse files- Dockerfile +33 -0
- README.md +327 -11
- requirements.txt +62 -0
- src/agents/base_agent.py +238 -0
- src/agents/config.py +111 -0
- src/agents/conversation_manager.py +275 -0
- src/agents/crypto_data/__init__.py +0 -0
- src/agents/crypto_data/agent.py +20 -0
- src/agents/crypto_data/config.py +17 -0
- src/agents/crypto_data/tools.py +436 -0
- src/agents/database/agent.py +109 -0
- src/agents/database/client.py +36 -0
- src/agents/database/config.py +14 -0
- src/agents/database/tools.py +179 -0
- src/agents/default/agent.py +13 -0
- src/agents/metadata.py +18 -0
- src/agents/supervisor/__init__.py +1 -0
- src/agents/supervisor/agent.py +306 -0
- src/agents/swap/agent.py +16 -0
- src/agents/swap/config.py +42 -0
- src/agents/swap/tools.py +49 -0
- src/app.py +197 -0
- src/models/chatMessage.py +152 -0
- src/routes/chat_manager_routes.py +90 -0
- src/service/chat_manager.py +280 -0
Dockerfile
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# syntax=docker/dockerfile:1
|
| 2 |
+
|
| 3 |
+
FROM python:3.12-slim
|
| 4 |
+
|
| 5 |
+
ENV PYTHONDONTWRITEBYTECODE=1 \
|
| 6 |
+
PYTHONUNBUFFERED=1 \
|
| 7 |
+
PIP_NO_CACHE_DIR=1 \
|
| 8 |
+
PYTHONPATH=/app \
|
| 9 |
+
PORT=8000
|
| 10 |
+
|
| 11 |
+
# Runtime libs for scientific stack (e.g., scikit-learn)
|
| 12 |
+
RUN apt-get update && apt-get install -y --no-install-recommends \
|
| 13 |
+
libgomp1 \
|
| 14 |
+
&& rm -rf /var/lib/apt/lists/*
|
| 15 |
+
|
| 16 |
+
WORKDIR /app
|
| 17 |
+
|
| 18 |
+
# Install Python dependencies first for better layer caching
|
| 19 |
+
COPY requirements.txt /app/requirements.txt
|
| 20 |
+
RUN pip install --upgrade pip setuptools wheel && \
|
| 21 |
+
pip install -r /app/requirements.txt
|
| 22 |
+
|
| 23 |
+
# Copy application source
|
| 24 |
+
COPY src /app/src
|
| 25 |
+
|
| 26 |
+
# Use a non-root user
|
| 27 |
+
RUN useradd -m -u 10001 appuser && chown -R appuser:appuser /app
|
| 28 |
+
USER appuser
|
| 29 |
+
|
| 30 |
+
EXPOSE 8000
|
| 31 |
+
|
| 32 |
+
# GEMINI_API_KEY must be provided at runtime
|
| 33 |
+
CMD ["sh", "-c", "uvicorn src.app:app --host 0.0.0.0 --port ${PORT}"]
|
README.md
CHANGED
|
@@ -1,11 +1,327 @@
|
|
| 1 |
-
-
|
| 2 |
-
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
--
|
| 10 |
-
|
| 11 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Zico Multi-Agent System
|
| 2 |
+
|
| 3 |
+
A sophisticated multi-agent system built with LangGraph, FastAPI, and Google's Gemini AI for intelligent conversation routing and specialized task handling.
|
| 4 |
+
|
| 5 |
+
## 🚀 Features
|
| 6 |
+
|
| 7 |
+
- **Multi-Agent Architecture**: Intelligent routing between specialized agents
|
| 8 |
+
- **LangGraph Integration**: Stateful conversation flows with proper state management
|
| 9 |
+
- **Real-time Crypto Data**: Live cryptocurrency price and market data
|
| 10 |
+
- **Conversation Management**: Persistent conversation state and history
|
| 11 |
+
- **Performance Monitoring**: Agent performance metrics and analytics
|
| 12 |
+
- **RESTful API**: Clean FastAPI endpoints for easy integration
|
| 13 |
+
- **Extensible Design**: Easy to add new agents and capabilities
|
| 14 |
+
|
| 15 |
+
## 🏗️ Architecture
|
| 16 |
+
|
| 17 |
+
### Core Components
|
| 18 |
+
|
| 19 |
+
1. **Supervisor Agent**: Routes messages to appropriate specialized agents
|
| 20 |
+
2. **Crypto Data Agent**: Handles cryptocurrency-related queries
|
| 21 |
+
3. **General Agent**: Manages general conversation and queries
|
| 22 |
+
4. **Conversation Manager**: Manages conversation state and persistence
|
| 23 |
+
5. **Agent Registry**: Central registry for all agents in the system
|
| 24 |
+
|
| 25 |
+
### LangGraph Flow
|
| 26 |
+
|
| 27 |
+
```
|
| 28 |
+
User Message → Supervisor → Route Decision → Specialized Agent → Response
|
| 29 |
+
↓ ↓ ↓ ↓ ↓
|
| 30 |
+
Conversation → Context → State → Processing → Agent Response
|
| 31 |
+
```
|
| 32 |
+
|
| 33 |
+
## 📦 Installation
|
| 34 |
+
|
| 35 |
+
1. **Clone the repository**
|
| 36 |
+
```bash
|
| 37 |
+
git clone <repository-url>
|
| 38 |
+
cd new_zico
|
| 39 |
+
```
|
| 40 |
+
|
| 41 |
+
2. **Create virtual environment**
|
| 42 |
+
```bash
|
| 43 |
+
python -m venv .venv
|
| 44 |
+
source .venv/bin/activate # On Windows: .venv\Scripts\activate
|
| 45 |
+
```
|
| 46 |
+
|
| 47 |
+
3. **Install dependencies**
|
| 48 |
+
```bash
|
| 49 |
+
pip install -r requirements.txt
|
| 50 |
+
```
|
| 51 |
+
|
| 52 |
+
4. **Set up environment variables**
|
| 53 |
+
```bash
|
| 54 |
+
cp .env.example .env
|
| 55 |
+
# Edit .env and add your GEMINI_API_KEY
|
| 56 |
+
```
|
| 57 |
+
|
| 58 |
+
5. **Run the application**
|
| 59 |
+
```bash
|
| 60 |
+
python -m uvicorn src.app:app --reload --host 0.0.0.0 --port 8000
|
| 61 |
+
```
|
| 62 |
+
|
| 63 |
+
## 🔧 Configuration
|
| 64 |
+
|
| 65 |
+
### Environment Variables
|
| 66 |
+
|
| 67 |
+
```env
|
| 68 |
+
GEMINI_API_KEY=your_gemini_api_key_here
|
| 69 |
+
GEMINI_MODEL=gemini-1.5-pro
|
| 70 |
+
GEMINI_EMBEDDING_MODEL=models/embedding-001
|
| 71 |
+
```
|
| 72 |
+
|
| 73 |
+
### Agent Configuration
|
| 74 |
+
|
| 75 |
+
Agents can be configured in `src/agents/config.py`:
|
| 76 |
+
|
| 77 |
+
```python
|
| 78 |
+
AGENTS_CONFIG = {
|
| 79 |
+
"agents": [
|
| 80 |
+
{
|
| 81 |
+
"name": "crypto_data",
|
| 82 |
+
"description": "Handles cryptocurrency-related queries",
|
| 83 |
+
"type": "specialized",
|
| 84 |
+
"enabled": True,
|
| 85 |
+
"priority": 1
|
| 86 |
+
},
|
| 87 |
+
{
|
| 88 |
+
"name": "general",
|
| 89 |
+
"description": "Handles general conversation and queries",
|
| 90 |
+
"type": "general",
|
| 91 |
+
"enabled": True,
|
| 92 |
+
"priority": 2
|
| 93 |
+
}
|
| 94 |
+
]
|
| 95 |
+
}
|
| 96 |
+
```
|
| 97 |
+
|
| 98 |
+
## 📡 API Usage
|
| 99 |
+
|
| 100 |
+
### Chat Endpoint
|
| 101 |
+
|
| 102 |
+
```bash
|
| 103 |
+
curl -X POST "http://localhost:8000/chat" \
|
| 104 |
+
-H "Content-Type: application/json" \
|
| 105 |
+
-d '{
|
| 106 |
+
"message": "What is the price of Bitcoin?",
|
| 107 |
+
"user_id": "user123",
|
| 108 |
+
"conversation_id": "conv456"
|
| 109 |
+
}'
|
| 110 |
+
```
|
| 111 |
+
|
| 112 |
+
**Response:**
|
| 113 |
+
```json
|
| 114 |
+
{
|
| 115 |
+
"response": "The current price of Bitcoin is $43,250.50",
|
| 116 |
+
"agent_name": "crypto_agent",
|
| 117 |
+
"agent_type": "crypto_data",
|
| 118 |
+
"conversation_id": "conv456",
|
| 119 |
+
"message_id": "msg789",
|
| 120 |
+
"metadata": {
|
| 121 |
+
"price_source": "CoinGecko",
|
| 122 |
+
"timestamp": "2024-01-15T10:30:00Z"
|
| 123 |
+
},
|
| 124 |
+
"next_agent": null,
|
| 125 |
+
"requires_followup": false,
|
| 126 |
+
"timestamp": "2024-01-15T10:30:00Z"
|
| 127 |
+
}
|
| 128 |
+
```
|
| 129 |
+
|
| 130 |
+
### Conversation Management
|
| 131 |
+
|
| 132 |
+
```bash
|
| 133 |
+
# Get all conversations for a user
|
| 134 |
+
GET /conversations/{user_id}
|
| 135 |
+
|
| 136 |
+
# Get messages from a specific conversation
|
| 137 |
+
GET /conversations/{user_id}/{conversation_id}/messages
|
| 138 |
+
|
| 139 |
+
# Delete a conversation
|
| 140 |
+
DELETE /conversations/{user_id}/{conversation_id}
|
| 141 |
+
|
| 142 |
+
# Reset a conversation (clear messages)
|
| 143 |
+
POST /conversations/{user_id}/reset
|
| 144 |
+
```
|
| 145 |
+
|
| 146 |
+
## 🤖 Adding New Agents
|
| 147 |
+
|
| 148 |
+
### 1. Create Agent Class
|
| 149 |
+
|
| 150 |
+
```python
|
| 151 |
+
# src/agents/my_agent/agent.py
|
| 152 |
+
from src.agents.base_agent import BaseAgent
|
| 153 |
+
from src.models.chatMessage import AgentType, AgentResponse
|
| 154 |
+
|
| 155 |
+
class MyAgent(BaseAgent):
|
| 156 |
+
def __init__(self, llm):
|
| 157 |
+
super().__init__(
|
| 158 |
+
name="my_agent",
|
| 159 |
+
agent_type=AgentType.SPECIALIZED,
|
| 160 |
+
llm=llm,
|
| 161 |
+
description="Handles specific tasks"
|
| 162 |
+
)
|
| 163 |
+
|
| 164 |
+
async def process_message(self, message: str, context: Dict[str, Any] = None) -> AgentResponse:
|
| 165 |
+
# Your agent logic here
|
| 166 |
+
response_content = "Processed by MyAgent"
|
| 167 |
+
return self.create_agent_response(response_content)
|
| 168 |
+
|
| 169 |
+
def get_capabilities(self) -> List[str]:
|
| 170 |
+
return ["task1", "task2", "task3"]
|
| 171 |
+
|
| 172 |
+
def can_handle(self, message: str, context: Dict[str, Any] = None) -> bool:
|
| 173 |
+
# Define when this agent should handle messages
|
| 174 |
+
return "my_keyword" in message.lower()
|
| 175 |
+
|
| 176 |
+
def get_confidence_score(self, message: str, context: Dict[str, Any] = None) -> float:
|
| 177 |
+
# Return confidence score (0.0 to 1.0)
|
| 178 |
+
if "my_keyword" in message.lower():
|
| 179 |
+
return 0.9
|
| 180 |
+
return 0.1
|
| 181 |
+
```
|
| 182 |
+
|
| 183 |
+
### 2. Register Agent
|
| 184 |
+
|
| 185 |
+
```python
|
| 186 |
+
# In your supervisor or main application
|
| 187 |
+
from src.agents.my_agent.agent import MyAgent
|
| 188 |
+
from src.agents.base_agent import agent_registry
|
| 189 |
+
|
| 190 |
+
my_agent = MyAgent(llm)
|
| 191 |
+
agent_registry.register_agent(my_agent)
|
| 192 |
+
```
|
| 193 |
+
|
| 194 |
+
### 3. Update Supervisor Routing
|
| 195 |
+
|
| 196 |
+
Add routing logic in the supervisor's `_route_to_agent` method:
|
| 197 |
+
|
| 198 |
+
```python
|
| 199 |
+
def _route_to_agent(self, state: ConversationState) -> str:
|
| 200 |
+
# ... existing logic ...
|
| 201 |
+
|
| 202 |
+
# Add your agent routing
|
| 203 |
+
if "my_keyword" in content:
|
| 204 |
+
return "my_agent"
|
| 205 |
+
|
| 206 |
+
# ... rest of logic ...
|
| 207 |
+
```
|
| 208 |
+
|
| 209 |
+
## 📊 Monitoring and Analytics
|
| 210 |
+
|
| 211 |
+
### Agent Performance
|
| 212 |
+
|
| 213 |
+
```bash
|
| 214 |
+
# Get performance metrics for all agents
|
| 215 |
+
GET /agents/performance
|
| 216 |
+
|
| 217 |
+
# Get specific agent info
|
| 218 |
+
GET /agents/{agent_name}/info
|
| 219 |
+
```
|
| 220 |
+
|
| 221 |
+
### Conversation Analytics
|
| 222 |
+
|
| 223 |
+
```bash
|
| 224 |
+
# Get conversation statistics
|
| 225 |
+
GET /conversations/{user_id}/stats
|
| 226 |
+
|
| 227 |
+
# Export conversation data
|
| 228 |
+
GET /conversations/{user_id}/{conversation_id}/export
|
| 229 |
+
```
|
| 230 |
+
|
| 231 |
+
## 🧪 Testing
|
| 232 |
+
|
| 233 |
+
Run the test suite:
|
| 234 |
+
|
| 235 |
+
```bash
|
| 236 |
+
# Run all tests
|
| 237 |
+
pytest
|
| 238 |
+
|
| 239 |
+
# Run with coverage
|
| 240 |
+
pytest --cov=src
|
| 241 |
+
|
| 242 |
+
# Run specific test file
|
| 243 |
+
pytest tests/test_supervisor.py
|
| 244 |
+
```
|
| 245 |
+
|
| 246 |
+
## 🚀 Production Deployment
|
| 247 |
+
|
| 248 |
+
### Docker Deployment
|
| 249 |
+
|
| 250 |
+
```dockerfile
|
| 251 |
+
FROM python:3.11-slim
|
| 252 |
+
|
| 253 |
+
WORKDIR /app
|
| 254 |
+
COPY requirements.txt .
|
| 255 |
+
RUN pip install -r requirements.txt
|
| 256 |
+
|
| 257 |
+
COPY . .
|
| 258 |
+
EXPOSE 8000
|
| 259 |
+
|
| 260 |
+
CMD ["uvicorn", "src.app:app", "--host", "0.0.0.0", "--port", "8000"]
|
| 261 |
+
```
|
| 262 |
+
|
| 263 |
+
### Environment Setup
|
| 264 |
+
|
| 265 |
+
For production, consider:
|
| 266 |
+
|
| 267 |
+
1. **Database**: Use PostgreSQL or MongoDB for conversation persistence
|
| 268 |
+
2. **Caching**: Redis for session management and caching
|
| 269 |
+
3. **Monitoring**: Prometheus + Grafana for metrics
|
| 270 |
+
4. **Logging**: Structured logging with ELK stack
|
| 271 |
+
5. **Security**: API keys, rate limiting, CORS configuration
|
| 272 |
+
|
| 273 |
+
## 🔍 Debugging
|
| 274 |
+
|
| 275 |
+
### Enable Debug Logging
|
| 276 |
+
|
| 277 |
+
```python
|
| 278 |
+
import logging
|
| 279 |
+
logging.basicConfig(level=logging.DEBUG)
|
| 280 |
+
```
|
| 281 |
+
|
| 282 |
+
### LangGraph Visualization
|
| 283 |
+
|
| 284 |
+
```python
|
| 285 |
+
# Visualize the conversation flow
|
| 286 |
+
from langgraph.graph import StateGraph
|
| 287 |
+
graph = supervisor.supervisor_graph
|
| 288 |
+
graph.get_graph().draw_mermaid()
|
| 289 |
+
```
|
| 290 |
+
|
| 291 |
+
## 📈 Performance Optimization
|
| 292 |
+
|
| 293 |
+
1. **Agent Caching**: Cache agent responses for similar queries
|
| 294 |
+
2. **Context Window**: Limit context messages to prevent token overflow
|
| 295 |
+
3. **Async Processing**: Use background tasks for non-critical operations
|
| 296 |
+
4. **Connection Pooling**: Reuse LLM connections
|
| 297 |
+
5. **Response Streaming**: Stream responses for long-running operations
|
| 298 |
+
|
| 299 |
+
## 🤝 Contributing
|
| 300 |
+
|
| 301 |
+
1. Fork the repository
|
| 302 |
+
2. Create a feature branch
|
| 303 |
+
3. Add tests for new functionality
|
| 304 |
+
4. Ensure all tests pass
|
| 305 |
+
5. Submit a pull request
|
| 306 |
+
|
| 307 |
+
## 📄 License
|
| 308 |
+
|
| 309 |
+
This project is licensed under the MIT License - see the LICENSE file for details.
|
| 310 |
+
|
| 311 |
+
## 🆘 Support
|
| 312 |
+
|
| 313 |
+
For support and questions:
|
| 314 |
+
|
| 315 |
+
- Create an issue in the repository
|
| 316 |
+
- Check the documentation
|
| 317 |
+
- Review the test examples
|
| 318 |
+
|
| 319 |
+
## 🔮 Roadmap
|
| 320 |
+
|
| 321 |
+
- [ ] Add more specialized agents (research, analysis, etc.)
|
| 322 |
+
- [ ] Implement agent learning and adaptation
|
| 323 |
+
- [ ] Add support for file uploads and document processing
|
| 324 |
+
- [ ] Implement real-time collaboration features
|
| 325 |
+
- [ ] Add advanced analytics and reporting
|
| 326 |
+
- [ ] Support for multiple LLM providers
|
| 327 |
+
- [ ] Mobile app integration
|
requirements.txt
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
python-dotenv==1.0.0
|
| 2 |
+
google-generativeai>=0.8.0
|
| 3 |
+
langchain-google-genai>=2.0.0
|
| 4 |
+
langchain-core>=0.2.43
|
| 5 |
+
fastapi==0.109.2
|
| 6 |
+
uvicorn==0.27.1
|
| 7 |
+
pydantic>=2.7.4
|
| 8 |
+
python-multipart==0.0.9
|
| 9 |
+
|
| 10 |
+
# LangGraph and multi-agent dependencies
|
| 11 |
+
langgraph>=0.2.70
|
| 12 |
+
langgraph-supervisor>=0.0.1
|
| 13 |
+
langchain>=0.2.0
|
| 14 |
+
langchain-community>=0.2.0
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
# Testing dependencies
|
| 18 |
+
pytest==8.0.0
|
| 19 |
+
pytest-cov==4.1.0
|
| 20 |
+
pytest-asyncio==0.23.5
|
| 21 |
+
|
| 22 |
+
# Additional dependencies
|
| 23 |
+
scikit-learn==1.4.0
|
| 24 |
+
numpy>=1.24.0
|
| 25 |
+
pandas>=2.0.0
|
| 26 |
+
|
| 27 |
+
# Monitoring and logging
|
| 28 |
+
structlog>=23.0.0
|
| 29 |
+
prometheus-client>=0.17.0
|
| 30 |
+
langsmith>=0.1.0
|
| 31 |
+
|
| 32 |
+
# Database (optional for production)
|
| 33 |
+
sqlalchemy>=2.0.0
|
| 34 |
+
alembic>=1.12.0
|
| 35 |
+
|
| 36 |
+
# Caching
|
| 37 |
+
redis>=5.0.0
|
| 38 |
+
|
| 39 |
+
# HTTP client
|
| 40 |
+
httpx>=0.25.0
|
| 41 |
+
requests>=2.31.0
|
| 42 |
+
|
| 43 |
+
# Data validation and serialization
|
| 44 |
+
marshmallow>=3.20.0
|
| 45 |
+
jsonschema>=4.19.0
|
| 46 |
+
|
| 47 |
+
# Async support
|
| 48 |
+
asyncio-mqtt>=0.16.0
|
| 49 |
+
aiofiles>=23.0.0
|
| 50 |
+
|
| 51 |
+
# Security
|
| 52 |
+
cryptography>=41.0.0
|
| 53 |
+
passlib>=1.7.4
|
| 54 |
+
|
| 55 |
+
# Development tools
|
| 56 |
+
black>=23.0.0
|
| 57 |
+
flake8>=6.0.0
|
| 58 |
+
mypy>=1.5.0
|
| 59 |
+
|
| 60 |
+
# ClickHouse dependencies
|
| 61 |
+
clickhouse-connect>=0.7.0
|
| 62 |
+
clickhouse-sqlalchemy==0.3.2
|
src/agents/base_agent.py
ADDED
|
@@ -0,0 +1,238 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from abc import ABC, abstractmethod
|
| 2 |
+
from typing import Dict, Any, List, Optional
|
| 3 |
+
import logging
|
| 4 |
+
from datetime import datetime
|
| 5 |
+
|
| 6 |
+
from langchain_core.messages import HumanMessage, SystemMessage, AIMessage
|
| 7 |
+
from langchain_core.language_models import BaseLanguageModel
|
| 8 |
+
|
| 9 |
+
from src.models.chatMessage import ChatMessage, AgentResponse, AgentType, MessageRole
|
| 10 |
+
from src.agents.config import Config
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class BaseAgent(ABC):
|
| 14 |
+
"""Base class for all agents in the multi-agent system"""
|
| 15 |
+
|
| 16 |
+
def __init__(self, name: str, agent_type: AgentType, llm: BaseLanguageModel, description: str = ""):
|
| 17 |
+
self.name = name
|
| 18 |
+
self.agent_type = agent_type
|
| 19 |
+
self.llm = llm
|
| 20 |
+
self.description = description
|
| 21 |
+
self.logger = logging.getLogger(f"{__name__}.{self.__class__.__name__}")
|
| 22 |
+
|
| 23 |
+
# Agent state
|
| 24 |
+
self.is_active = True
|
| 25 |
+
self.created_at = datetime.utcnow()
|
| 26 |
+
self.last_used = None
|
| 27 |
+
self.usage_count = 0
|
| 28 |
+
|
| 29 |
+
# Performance metrics
|
| 30 |
+
self.response_times = []
|
| 31 |
+
self.success_count = 0
|
| 32 |
+
self.error_count = 0
|
| 33 |
+
|
| 34 |
+
self.logger.info(f"Initialized agent {name} of type {agent_type}")
|
| 35 |
+
|
| 36 |
+
@abstractmethod
|
| 37 |
+
async def process_message(self, message: str, context: Dict[str, Any] = None) -> AgentResponse:
|
| 38 |
+
"""Process a message and return a response"""
|
| 39 |
+
pass
|
| 40 |
+
|
| 41 |
+
@abstractmethod
|
| 42 |
+
def get_capabilities(self) -> List[str]:
|
| 43 |
+
"""Return list of agent capabilities"""
|
| 44 |
+
pass
|
| 45 |
+
|
| 46 |
+
def can_handle(self, message: str, context: Dict[str, Any] = None) -> bool:
|
| 47 |
+
"""Determine if this agent can handle the given message"""
|
| 48 |
+
# Default implementation - can be overridden by subclasses
|
| 49 |
+
return True
|
| 50 |
+
|
| 51 |
+
def get_confidence_score(self, message: str, context: Dict[str, Any] = None) -> float:
|
| 52 |
+
"""Get confidence score for handling this message (0.0 to 1.0)"""
|
| 53 |
+
# Default implementation - can be overridden by subclasses
|
| 54 |
+
return 0.5
|
| 55 |
+
|
| 56 |
+
def prepare_context_messages(self, context: Dict[str, Any] = None) -> List[SystemMessage]:
|
| 57 |
+
"""Prepare context messages for the LLM"""
|
| 58 |
+
context_messages = []
|
| 59 |
+
|
| 60 |
+
if context:
|
| 61 |
+
# Add relevant context information
|
| 62 |
+
if context.get("crypto_related"):
|
| 63 |
+
context_messages.append(SystemMessage(
|
| 64 |
+
content="This conversation involves cryptocurrency-related topics."
|
| 65 |
+
))
|
| 66 |
+
|
| 67 |
+
if context.get("user_preferences"):
|
| 68 |
+
context_messages.append(SystemMessage(
|
| 69 |
+
content=f"User preferences: {context['user_preferences']}"
|
| 70 |
+
))
|
| 71 |
+
|
| 72 |
+
if context.get("conversation_history"):
|
| 73 |
+
context_messages.append(SystemMessage(
|
| 74 |
+
content=f"Previous context: {context['conversation_history']}"
|
| 75 |
+
))
|
| 76 |
+
|
| 77 |
+
return context_messages
|
| 78 |
+
|
| 79 |
+
def create_agent_response(
|
| 80 |
+
self,
|
| 81 |
+
content: str,
|
| 82 |
+
success: bool = True,
|
| 83 |
+
error_message: Optional[str] = None,
|
| 84 |
+
metadata: Dict[str, Any] = None,
|
| 85 |
+
tools_used: List[str] = None,
|
| 86 |
+
next_agent: Optional[str] = None,
|
| 87 |
+
requires_followup: bool = False
|
| 88 |
+
) -> AgentResponse:
|
| 89 |
+
"""Create a standardized agent response"""
|
| 90 |
+
return AgentResponse(
|
| 91 |
+
content=content,
|
| 92 |
+
agent_name=self.name,
|
| 93 |
+
agent_type=self.agent_type,
|
| 94 |
+
success=success,
|
| 95 |
+
error_message=error_message,
|
| 96 |
+
metadata=metadata or {},
|
| 97 |
+
tools_used=tools_used or [],
|
| 98 |
+
next_agent=next_agent,
|
| 99 |
+
requires_followup=requires_followup,
|
| 100 |
+
timestamp=datetime.utcnow()
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
def update_metrics(self, response_time: float, success: bool):
|
| 104 |
+
"""Update agent performance metrics"""
|
| 105 |
+
self.response_times.append(response_time)
|
| 106 |
+
self.last_used = datetime.utcnow()
|
| 107 |
+
self.usage_count += 1
|
| 108 |
+
|
| 109 |
+
if success:
|
| 110 |
+
self.success_count += 1
|
| 111 |
+
else:
|
| 112 |
+
self.error_count += 1
|
| 113 |
+
|
| 114 |
+
# Keep only last 100 response times
|
| 115 |
+
if len(self.response_times) > 100:
|
| 116 |
+
self.response_times = self.response_times[-100:]
|
| 117 |
+
|
| 118 |
+
def get_performance_metrics(self) -> Dict[str, Any]:
|
| 119 |
+
"""Get agent performance metrics"""
|
| 120 |
+
avg_response_time = sum(self.response_times) / len(self.response_times) if self.response_times else 0
|
| 121 |
+
success_rate = self.success_count / self.usage_count if self.usage_count > 0 else 0
|
| 122 |
+
|
| 123 |
+
return {
|
| 124 |
+
"name": self.name,
|
| 125 |
+
"agent_type": self.agent_type.value,
|
| 126 |
+
"usage_count": self.usage_count,
|
| 127 |
+
"success_count": self.success_count,
|
| 128 |
+
"error_count": self.error_count,
|
| 129 |
+
"success_rate": success_rate,
|
| 130 |
+
"average_response_time": avg_response_time,
|
| 131 |
+
"last_used": self.last_used.isoformat() if self.last_used else None,
|
| 132 |
+
"is_active": self.is_active
|
| 133 |
+
}
|
| 134 |
+
|
| 135 |
+
def activate(self):
|
| 136 |
+
"""Activate the agent"""
|
| 137 |
+
self.is_active = True
|
| 138 |
+
self.logger.info(f"Agent {self.name} activated")
|
| 139 |
+
|
| 140 |
+
def deactivate(self):
|
| 141 |
+
"""Deactivate the agent"""
|
| 142 |
+
self.is_active = False
|
| 143 |
+
self.logger.info(f"Agent {self.name} deactivated")
|
| 144 |
+
|
| 145 |
+
def reset_metrics(self):
|
| 146 |
+
"""Reset performance metrics"""
|
| 147 |
+
self.response_times = []
|
| 148 |
+
self.success_count = 0
|
| 149 |
+
self.error_count = 0
|
| 150 |
+
self.usage_count = 0
|
| 151 |
+
self.logger.info(f"Reset metrics for agent {self.name}")
|
| 152 |
+
|
| 153 |
+
def get_agent_info(self) -> Dict[str, Any]:
|
| 154 |
+
"""Get comprehensive agent information"""
|
| 155 |
+
return {
|
| 156 |
+
"name": self.name,
|
| 157 |
+
"type": self.agent_type.value,
|
| 158 |
+
"description": self.description,
|
| 159 |
+
"capabilities": self.get_capabilities(),
|
| 160 |
+
"is_active": self.is_active,
|
| 161 |
+
"created_at": self.created_at.isoformat(),
|
| 162 |
+
"last_used": self.last_used.isoformat() if self.last_used else None,
|
| 163 |
+
"performance_metrics": self.get_performance_metrics()
|
| 164 |
+
}
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
class AgentRegistry:
|
| 168 |
+
"""Registry for managing all agents in the system"""
|
| 169 |
+
|
| 170 |
+
def __init__(self):
|
| 171 |
+
self.agents: Dict[str, BaseAgent] = {}
|
| 172 |
+
self.logger = logging.getLogger(__name__)
|
| 173 |
+
|
| 174 |
+
def register_agent(self, agent: BaseAgent) -> None:
|
| 175 |
+
"""Register an agent"""
|
| 176 |
+
if agent.name in self.agents:
|
| 177 |
+
self.logger.warning(f"Agent {agent.name} already registered, overwriting")
|
| 178 |
+
|
| 179 |
+
self.agents[agent.name] = agent
|
| 180 |
+
self.logger.info(f"Registered agent {agent.name}")
|
| 181 |
+
|
| 182 |
+
def unregister_agent(self, agent_name: str) -> bool:
|
| 183 |
+
"""Unregister an agent"""
|
| 184 |
+
if agent_name in self.agents:
|
| 185 |
+
del self.agents[agent_name]
|
| 186 |
+
self.logger.info(f"Unregistered agent {agent_name}")
|
| 187 |
+
return True
|
| 188 |
+
return False
|
| 189 |
+
|
| 190 |
+
def get_agent(self, agent_name: str) -> Optional[BaseAgent]:
|
| 191 |
+
"""Get agent by name"""
|
| 192 |
+
return self.agents.get(agent_name)
|
| 193 |
+
|
| 194 |
+
def get_active_agents(self) -> List[BaseAgent]:
|
| 195 |
+
"""Get all active agents"""
|
| 196 |
+
return [agent for agent in self.agents.values() if agent.is_active]
|
| 197 |
+
|
| 198 |
+
def get_agents_by_type(self, agent_type: AgentType) -> List[BaseAgent]:
|
| 199 |
+
"""Get agents by type"""
|
| 200 |
+
return [agent for agent in self.agents.values() if agent.agent_type == agent_type]
|
| 201 |
+
|
| 202 |
+
def find_best_agent(self, message: str, context: Dict[str, Any] = None) -> Optional[BaseAgent]:
|
| 203 |
+
"""Find the best agent to handle a message"""
|
| 204 |
+
best_agent = None
|
| 205 |
+
best_score = 0.0
|
| 206 |
+
|
| 207 |
+
for agent in self.get_active_agents():
|
| 208 |
+
if agent.can_handle(message, context):
|
| 209 |
+
confidence = agent.get_confidence_score(message, context)
|
| 210 |
+
if confidence > best_score:
|
| 211 |
+
best_score = confidence
|
| 212 |
+
best_agent = agent
|
| 213 |
+
|
| 214 |
+
return best_agent
|
| 215 |
+
|
| 216 |
+
def get_all_agents_info(self) -> List[Dict[str, Any]]:
|
| 217 |
+
"""Get information about all agents"""
|
| 218 |
+
return [agent.get_agent_info() for agent in self.agents.values()]
|
| 219 |
+
|
| 220 |
+
def get_agent_performance_summary(self) -> Dict[str, Any]:
|
| 221 |
+
"""Get performance summary for all agents"""
|
| 222 |
+
total_agents = len(self.agents)
|
| 223 |
+
active_agents = len(self.get_active_agents())
|
| 224 |
+
total_usage = sum(agent.usage_count for agent in self.agents.values())
|
| 225 |
+
total_success = sum(agent.success_count for agent in self.agents.values())
|
| 226 |
+
|
| 227 |
+
return {
|
| 228 |
+
"total_agents": total_agents,
|
| 229 |
+
"active_agents": active_agents,
|
| 230 |
+
"total_usage": total_usage,
|
| 231 |
+
"total_success": total_success,
|
| 232 |
+
"overall_success_rate": total_success / total_usage if total_usage > 0 else 0,
|
| 233 |
+
"agents": self.get_all_agents_info()
|
| 234 |
+
}
|
| 235 |
+
|
| 236 |
+
|
| 237 |
+
# Global agent registry
|
| 238 |
+
agent_registry = AgentRegistry()
|
src/agents/config.py
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from dotenv import load_dotenv
|
| 3 |
+
from langchain_google_genai import ChatGoogleGenerativeAI, GoogleGenerativeAIEmbeddings
|
| 4 |
+
from typing import Optional
|
| 5 |
+
|
| 6 |
+
load_dotenv()
|
| 7 |
+
|
| 8 |
+
gemini_api_key = os.getenv("GEMINI_API_KEY")
|
| 9 |
+
if not gemini_api_key:
|
| 10 |
+
raise ValueError("GEMINI_API_KEY não encontrada nas variáveis de ambiente")
|
| 11 |
+
|
| 12 |
+
class Config:
|
| 13 |
+
# Model configuration
|
| 14 |
+
GEMINI_MODEL = "gemini-1.5-pro"
|
| 15 |
+
GEMINI_EMBEDDING_MODEL = "models/embedding-001"
|
| 16 |
+
GEMINI_API_KEY = gemini_api_key
|
| 17 |
+
|
| 18 |
+
# Application configuration
|
| 19 |
+
MAX_UPLOAD_LENGTH = 16 * 1024 * 1024
|
| 20 |
+
MAX_CONVERSATION_LENGTH = 100 # Maximum messages per conversation
|
| 21 |
+
MAX_CONTEXT_MESSAGES = 10 # Maximum messages to include in context
|
| 22 |
+
|
| 23 |
+
# Agent configuration
|
| 24 |
+
AGENTS_CONFIG = {
|
| 25 |
+
"agents": [
|
| 26 |
+
{
|
| 27 |
+
"name": "crypto_data",
|
| 28 |
+
"description": "Handles cryptocurrency-related queries",
|
| 29 |
+
"type": "specialized",
|
| 30 |
+
"enabled": True,
|
| 31 |
+
"priority": 1
|
| 32 |
+
},
|
| 33 |
+
{
|
| 34 |
+
"name": "general",
|
| 35 |
+
"description": "Handles general conversation and queries",
|
| 36 |
+
"type": "general",
|
| 37 |
+
"enabled": True,
|
| 38 |
+
"priority": 2
|
| 39 |
+
}
|
| 40 |
+
]
|
| 41 |
+
}
|
| 42 |
+
|
| 43 |
+
# LangGraph configuration
|
| 44 |
+
LANGGRAPH_CONFIG = {
|
| 45 |
+
"max_iterations": 10,
|
| 46 |
+
"timeout": 30,
|
| 47 |
+
"memory_window": 10,
|
| 48 |
+
"enable_memory": True
|
| 49 |
+
}
|
| 50 |
+
|
| 51 |
+
# Conversation configuration
|
| 52 |
+
CONVERSATION_CONFIG = {
|
| 53 |
+
"default_user_id": "anonymous",
|
| 54 |
+
"max_conversations_per_user": 50,
|
| 55 |
+
"conversation_timeout_hours": 24,
|
| 56 |
+
"enable_context_extraction": True
|
| 57 |
+
}
|
| 58 |
+
|
| 59 |
+
# LLM instances (singleton pattern)
|
| 60 |
+
_llm_instance: Optional[ChatGoogleGenerativeAI] = None
|
| 61 |
+
_embeddings_instance: Optional[GoogleGenerativeAIEmbeddings] = None
|
| 62 |
+
|
| 63 |
+
@classmethod
|
| 64 |
+
def get_llm(cls) -> ChatGoogleGenerativeAI:
|
| 65 |
+
"""Get or create LLM instance (singleton)"""
|
| 66 |
+
if cls._llm_instance is None:
|
| 67 |
+
cls._llm_instance = ChatGoogleGenerativeAI(
|
| 68 |
+
model=cls.GEMINI_MODEL,
|
| 69 |
+
temperature=0.7,
|
| 70 |
+
google_api_key=cls.GEMINI_API_KEY
|
| 71 |
+
)
|
| 72 |
+
return cls._llm_instance
|
| 73 |
+
|
| 74 |
+
@classmethod
|
| 75 |
+
def get_embeddings(cls) -> GoogleGenerativeAIEmbeddings:
|
| 76 |
+
"""Get or create embeddings instance (singleton)"""
|
| 77 |
+
if cls._embeddings_instance is None:
|
| 78 |
+
cls._embeddings_instance = GoogleGenerativeAIEmbeddings(
|
| 79 |
+
model=cls.GEMINI_EMBEDDING_MODEL,
|
| 80 |
+
google_api_key=cls.GEMINI_API_KEY
|
| 81 |
+
)
|
| 82 |
+
return cls._embeddings_instance
|
| 83 |
+
|
| 84 |
+
@classmethod
|
| 85 |
+
def get_agent_config(cls, agent_name: str) -> Optional[dict]:
|
| 86 |
+
"""Get configuration for a specific agent"""
|
| 87 |
+
for agent in cls.AGENTS_CONFIG["agents"]:
|
| 88 |
+
if agent["name"] == agent_name:
|
| 89 |
+
return agent
|
| 90 |
+
return None
|
| 91 |
+
|
| 92 |
+
@classmethod
|
| 93 |
+
def get_enabled_agents(cls) -> list:
|
| 94 |
+
"""Get list of enabled agents"""
|
| 95 |
+
return [
|
| 96 |
+
agent for agent in cls.AGENTS_CONFIG["agents"]
|
| 97 |
+
if agent.get("enabled", True)
|
| 98 |
+
]
|
| 99 |
+
|
| 100 |
+
@classmethod
|
| 101 |
+
def validate_config(cls) -> bool:
|
| 102 |
+
"""Validate configuration"""
|
| 103 |
+
try:
|
| 104 |
+
# Test LLM connection
|
| 105 |
+
llm = cls.get_llm()
|
| 106 |
+
# Test embeddings connection
|
| 107 |
+
embeddings = cls.get_embeddings()
|
| 108 |
+
return True
|
| 109 |
+
except Exception as e:
|
| 110 |
+
print(f"Configuration validation failed: {e}")
|
| 111 |
+
return False
|
src/agents/conversation_manager.py
ADDED
|
@@ -0,0 +1,275 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
import uuid
|
| 3 |
+
from typing import Dict, List, Optional, Any
|
| 4 |
+
from datetime import datetime, timedelta
|
| 5 |
+
from dataclasses import dataclass, asdict
|
| 6 |
+
import json
|
| 7 |
+
|
| 8 |
+
from src.models.chatMessage import ConversationState, ChatMessage, MessageRole, AgentType
|
| 9 |
+
|
| 10 |
+
logger = logging.getLogger(__name__)
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
@dataclass
|
| 14 |
+
class ConversationMetadata:
|
| 15 |
+
"""Metadata for conversation tracking"""
|
| 16 |
+
conversation_id: str
|
| 17 |
+
user_id: str
|
| 18 |
+
created_at: datetime
|
| 19 |
+
updated_at: datetime
|
| 20 |
+
message_count: int
|
| 21 |
+
current_agent: Optional[str]
|
| 22 |
+
is_active: bool
|
| 23 |
+
context_summary: Dict[str, Any]
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
class ConversationManager:
|
| 27 |
+
"""Manages conversation state and persistence for multi-agent system"""
|
| 28 |
+
|
| 29 |
+
def __init__(self):
|
| 30 |
+
self.conversations: Dict[str, ConversationState] = {}
|
| 31 |
+
self.metadata: Dict[str, ConversationMetadata] = {}
|
| 32 |
+
self.user_conversations: Dict[str, List[str]] = {}
|
| 33 |
+
|
| 34 |
+
def create_conversation(self, user_id: str, conversation_id: Optional[str] = None) -> str:
|
| 35 |
+
"""Create a new conversation"""
|
| 36 |
+
if not conversation_id:
|
| 37 |
+
conversation_id = str(uuid.uuid4())
|
| 38 |
+
|
| 39 |
+
key = f"{user_id}:{conversation_id}"
|
| 40 |
+
|
| 41 |
+
# Create conversation state
|
| 42 |
+
conversation_state = ConversationState(
|
| 43 |
+
conversation_id=conversation_id,
|
| 44 |
+
user_id=user_id,
|
| 45 |
+
messages=[],
|
| 46 |
+
context={},
|
| 47 |
+
memory={},
|
| 48 |
+
agent_history=[],
|
| 49 |
+
current_agent=None,
|
| 50 |
+
last_message_id=None,
|
| 51 |
+
created_at=datetime.utcnow(),
|
| 52 |
+
updated_at=datetime.utcnow(),
|
| 53 |
+
is_active=True
|
| 54 |
+
)
|
| 55 |
+
|
| 56 |
+
# Create metadata
|
| 57 |
+
metadata = ConversationMetadata(
|
| 58 |
+
conversation_id=conversation_id,
|
| 59 |
+
user_id=user_id,
|
| 60 |
+
created_at=datetime.utcnow(),
|
| 61 |
+
updated_at=datetime.utcnow(),
|
| 62 |
+
message_count=0,
|
| 63 |
+
current_agent=None,
|
| 64 |
+
is_active=True,
|
| 65 |
+
context_summary={}
|
| 66 |
+
)
|
| 67 |
+
|
| 68 |
+
# Store conversation
|
| 69 |
+
self.conversations[key] = conversation_state
|
| 70 |
+
self.metadata[key] = metadata
|
| 71 |
+
|
| 72 |
+
# Update user conversations
|
| 73 |
+
if user_id not in self.user_conversations:
|
| 74 |
+
self.user_conversations[user_id] = []
|
| 75 |
+
self.user_conversations[user_id].append(conversation_id)
|
| 76 |
+
|
| 77 |
+
logger.info(f"Created conversation {conversation_id} for user {user_id}")
|
| 78 |
+
return conversation_id
|
| 79 |
+
|
| 80 |
+
def get_conversation(self, conversation_id: str, user_id: str) -> Optional[ConversationState]:
|
| 81 |
+
"""Get conversation by ID and user"""
|
| 82 |
+
key = f"{user_id}:{conversation_id}"
|
| 83 |
+
return self.conversations.get(key)
|
| 84 |
+
|
| 85 |
+
def get_or_create_conversation(self, conversation_id: str, user_id: str) -> ConversationState:
|
| 86 |
+
"""Get existing conversation or create new one"""
|
| 87 |
+
conversation = self.get_conversation(conversation_id, user_id)
|
| 88 |
+
if not conversation:
|
| 89 |
+
self.create_conversation(user_id, conversation_id)
|
| 90 |
+
conversation = self.get_conversation(conversation_id, user_id)
|
| 91 |
+
return conversation
|
| 92 |
+
|
| 93 |
+
def add_message(self, conversation_id: str, user_id: str, message: ChatMessage) -> None:
|
| 94 |
+
"""Add message to conversation"""
|
| 95 |
+
conversation = self.get_or_create_conversation(conversation_id, user_id)
|
| 96 |
+
key = f"{user_id}:{conversation_id}"
|
| 97 |
+
|
| 98 |
+
# Add message
|
| 99 |
+
conversation.messages.append(message)
|
| 100 |
+
conversation.last_message_id = message.message_id
|
| 101 |
+
conversation.updated_at = datetime.utcnow()
|
| 102 |
+
|
| 103 |
+
# Update metadata
|
| 104 |
+
if key in self.metadata:
|
| 105 |
+
self.metadata[key].message_count = len(conversation.messages)
|
| 106 |
+
self.metadata[key].updated_at = datetime.utcnow()
|
| 107 |
+
self.metadata[key].current_agent = conversation.current_agent
|
| 108 |
+
|
| 109 |
+
logger.info(f"Added message to conversation {conversation_id}")
|
| 110 |
+
|
| 111 |
+
def update_conversation_context(self, conversation_id: str, user_id: str, context_updates: Dict[str, Any]) -> None:
|
| 112 |
+
"""Update conversation context"""
|
| 113 |
+
conversation = self.get_conversation(conversation_id, user_id)
|
| 114 |
+
if conversation:
|
| 115 |
+
conversation.context.update(context_updates)
|
| 116 |
+
conversation.updated_at = datetime.utcnow()
|
| 117 |
+
|
| 118 |
+
# Update metadata
|
| 119 |
+
key = f"{user_id}:{conversation_id}"
|
| 120 |
+
if key in self.metadata:
|
| 121 |
+
self.metadata[key].context_summary.update(context_updates)
|
| 122 |
+
self.metadata[key].updated_at = datetime.utcnow()
|
| 123 |
+
|
| 124 |
+
def update_agent_history(self, conversation_id: str, user_id: str, agent_info: Dict[str, Any]) -> None:
|
| 125 |
+
"""Update agent interaction history"""
|
| 126 |
+
conversation = self.get_conversation(conversation_id, user_id)
|
| 127 |
+
if conversation:
|
| 128 |
+
conversation.agent_history.append(agent_info)
|
| 129 |
+
conversation.updated_at = datetime.utcnow()
|
| 130 |
+
|
| 131 |
+
def get_conversation_messages(self, conversation_id: str, user_id: str, limit: Optional[int] = None) -> List[ChatMessage]:
|
| 132 |
+
"""Get messages from conversation"""
|
| 133 |
+
conversation = self.get_conversation(conversation_id, user_id)
|
| 134 |
+
if not conversation:
|
| 135 |
+
return []
|
| 136 |
+
|
| 137 |
+
messages = conversation.messages
|
| 138 |
+
if limit:
|
| 139 |
+
messages = messages[-limit:]
|
| 140 |
+
|
| 141 |
+
return messages
|
| 142 |
+
|
| 143 |
+
def get_user_conversations(self, user_id: str) -> List[Dict[str, Any]]:
|
| 144 |
+
"""Get all conversations for a user"""
|
| 145 |
+
user_conversations = []
|
| 146 |
+
|
| 147 |
+
for conversation_id in self.user_conversations.get(user_id, []):
|
| 148 |
+
key = f"{user_id}:{conversation_id}"
|
| 149 |
+
metadata = self.metadata.get(key)
|
| 150 |
+
|
| 151 |
+
if metadata:
|
| 152 |
+
user_conversations.append(asdict(metadata))
|
| 153 |
+
|
| 154 |
+
return user_conversations
|
| 155 |
+
|
| 156 |
+
def delete_conversation(self, conversation_id: str, user_id: str) -> bool:
|
| 157 |
+
"""Delete a conversation"""
|
| 158 |
+
key = f"{user_id}:{conversation_id}"
|
| 159 |
+
|
| 160 |
+
if key in self.conversations:
|
| 161 |
+
del self.conversations[key]
|
| 162 |
+
|
| 163 |
+
if key in self.metadata:
|
| 164 |
+
del self.metadata[key]
|
| 165 |
+
|
| 166 |
+
# Remove from user conversations
|
| 167 |
+
if user_id in self.user_conversations:
|
| 168 |
+
if conversation_id in self.user_conversations[user_id]:
|
| 169 |
+
self.user_conversations[user_id].remove(conversation_id)
|
| 170 |
+
|
| 171 |
+
logger.info(f"Deleted conversation {conversation_id} for user {user_id}")
|
| 172 |
+
return True
|
| 173 |
+
|
| 174 |
+
def reset_conversation(self, conversation_id: str, user_id: str) -> None:
|
| 175 |
+
"""Reset conversation (clear messages but keep conversation)"""
|
| 176 |
+
conversation = self.get_conversation(conversation_id, user_id)
|
| 177 |
+
if conversation:
|
| 178 |
+
conversation.messages = []
|
| 179 |
+
conversation.context = {}
|
| 180 |
+
conversation.agent_history = []
|
| 181 |
+
conversation.current_agent = None
|
| 182 |
+
conversation.last_message_id = None
|
| 183 |
+
conversation.updated_at = datetime.utcnow()
|
| 184 |
+
|
| 185 |
+
# Update metadata
|
| 186 |
+
key = f"{user_id}:{conversation_id}"
|
| 187 |
+
if key in self.metadata:
|
| 188 |
+
self.metadata[key].message_count = 0
|
| 189 |
+
self.metadata[key].current_agent = None
|
| 190 |
+
self.metadata[key].context_summary = {}
|
| 191 |
+
self.metadata[key].updated_at = datetime.utcnow()
|
| 192 |
+
|
| 193 |
+
def cleanup_old_conversations(self, max_age_hours: int = 24) -> int:
|
| 194 |
+
"""Clean up old conversations"""
|
| 195 |
+
cutoff_time = datetime.utcnow() - timedelta(hours=max_age_hours)
|
| 196 |
+
deleted_count = 0
|
| 197 |
+
|
| 198 |
+
conversations_to_delete = []
|
| 199 |
+
|
| 200 |
+
for key, metadata in self.metadata.items():
|
| 201 |
+
if metadata.updated_at < cutoff_time and not metadata.is_active:
|
| 202 |
+
conversations_to_delete.append(key)
|
| 203 |
+
|
| 204 |
+
for key in conversations_to_delete:
|
| 205 |
+
user_id, conversation_id = key.split(":", 1)
|
| 206 |
+
if self.delete_conversation(conversation_id, user_id):
|
| 207 |
+
deleted_count += 1
|
| 208 |
+
|
| 209 |
+
logger.info(f"Cleaned up {deleted_count} old conversations")
|
| 210 |
+
return deleted_count
|
| 211 |
+
|
| 212 |
+
def get_conversation_stats(self, user_id: str) -> Dict[str, Any]:
|
| 213 |
+
"""Get conversation statistics for a user"""
|
| 214 |
+
user_conversations = self.get_user_conversations(user_id)
|
| 215 |
+
|
| 216 |
+
total_conversations = len(user_conversations)
|
| 217 |
+
active_conversations = sum(1 for conv in user_conversations if conv["is_active"])
|
| 218 |
+
total_messages = sum(conv["message_count"] for conv in user_conversations)
|
| 219 |
+
|
| 220 |
+
# Agent usage statistics
|
| 221 |
+
agent_usage = {}
|
| 222 |
+
for conv in user_conversations:
|
| 223 |
+
conversation = self.get_conversation(conv["conversation_id"], user_id)
|
| 224 |
+
if conversation:
|
| 225 |
+
for agent_info in conversation.agent_history:
|
| 226 |
+
agent_name = agent_info.get("agent", "unknown")
|
| 227 |
+
agent_usage[agent_name] = agent_usage.get(agent_name, 0) + 1
|
| 228 |
+
|
| 229 |
+
return {
|
| 230 |
+
"total_conversations": total_conversations,
|
| 231 |
+
"active_conversations": active_conversations,
|
| 232 |
+
"total_messages": total_messages,
|
| 233 |
+
"agent_usage": agent_usage,
|
| 234 |
+
"average_messages_per_conversation": total_messages / total_conversations if total_conversations > 0 else 0
|
| 235 |
+
}
|
| 236 |
+
|
| 237 |
+
def export_conversation(self, conversation_id: str, user_id: str) -> Dict[str, Any]:
|
| 238 |
+
"""Export conversation data"""
|
| 239 |
+
conversation = self.get_conversation(conversation_id, user_id)
|
| 240 |
+
if not conversation:
|
| 241 |
+
return {}
|
| 242 |
+
|
| 243 |
+
return {
|
| 244 |
+
"conversation_id": conversation_id,
|
| 245 |
+
"user_id": user_id,
|
| 246 |
+
"messages": [msg.dict() for msg in conversation.messages],
|
| 247 |
+
"context": conversation.context,
|
| 248 |
+
"agent_history": conversation.agent_history,
|
| 249 |
+
"metadata": asdict(self.metadata.get(f"{user_id}:{conversation_id}", {}))
|
| 250 |
+
}
|
| 251 |
+
|
| 252 |
+
def import_conversation(self, conversation_data: Dict[str, Any]) -> str:
|
| 253 |
+
"""Import conversation data"""
|
| 254 |
+
conversation_id = conversation_data.get("conversation_id", str(uuid.uuid4()))
|
| 255 |
+
user_id = conversation_data.get("user_id", "anonymous")
|
| 256 |
+
|
| 257 |
+
# Create conversation
|
| 258 |
+
self.create_conversation(user_id, conversation_id)
|
| 259 |
+
|
| 260 |
+
# Import messages
|
| 261 |
+
for msg_data in conversation_data.get("messages", []):
|
| 262 |
+
message = ChatMessage(**msg_data)
|
| 263 |
+
self.add_message(conversation_id, user_id, message)
|
| 264 |
+
|
| 265 |
+
# Import context and history
|
| 266 |
+
conversation = self.get_conversation(conversation_id, user_id)
|
| 267 |
+
if conversation:
|
| 268 |
+
conversation.context.update(conversation_data.get("context", {}))
|
| 269 |
+
conversation.agent_history.extend(conversation_data.get("agent_history", []))
|
| 270 |
+
|
| 271 |
+
return conversation_id
|
| 272 |
+
|
| 273 |
+
|
| 274 |
+
# Global conversation manager instance
|
| 275 |
+
conversation_manager = ConversationManager()
|
src/agents/crypto_data/__init__.py
ADDED
|
File without changes
|
src/agents/crypto_data/agent.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from src.agents.crypto_data.tools import get_tools
|
| 3 |
+
from langgraph.prebuilt import create_react_agent
|
| 4 |
+
|
| 5 |
+
logger = logging.getLogger(__name__)
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class CryptoDataAgent():
|
| 9 |
+
"""Agent for handling cryptocurrency-related queries and data retrieval."""
|
| 10 |
+
|
| 11 |
+
def __init__(self, llm):
|
| 12 |
+
self.llm = llm
|
| 13 |
+
|
| 14 |
+
self.agent = create_react_agent(
|
| 15 |
+
model=llm,
|
| 16 |
+
tools=get_tools(),
|
| 17 |
+
name="crypto_agent"
|
| 18 |
+
)
|
| 19 |
+
|
| 20 |
+
|
src/agents/crypto_data/config.py
ADDED
|
@@ -0,0 +1,17 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
class Config:
|
| 2 |
+
|
| 3 |
+
# API endpoints
|
| 4 |
+
COINGECKO_BASE_URL = "https://api.coingecko.com/api/v3"
|
| 5 |
+
DEFILLAMA_BASE_URL = "https://api.llama.fi"
|
| 6 |
+
PRICE_SUCCESS_MESSAGE = "The price of {coin_name} is ${price:,}"
|
| 7 |
+
PRICE_FAILURE_MESSAGE = "Failed to retrieve price. Please enter a valid coin name."
|
| 8 |
+
FLOOR_PRICE_SUCCESS_MESSAGE = "The floor price of {nft_name} is ${floor_price:,}"
|
| 9 |
+
FLOOR_PRICE_FAILURE_MESSAGE = "Failed to retrieve floor price. Please enter a valid NFT name."
|
| 10 |
+
TVL_SUCCESS_MESSAGE = "The TVL of {protocol_name} is ${tvl:,}"
|
| 11 |
+
TVL_FAILURE_MESSAGE = "Failed to retrieve TVL. Please enter a valid protocol name."
|
| 12 |
+
FDV_SUCCESS_MESSAGE = "The fully diluted valuation of {coin_name} is ${fdv:,}"
|
| 13 |
+
FDV_FAILURE_MESSAGE = "Failed to retrieve FDV. Please enter a valid coin name."
|
| 14 |
+
MARKET_CAP_SUCCESS_MESSAGE = "The market cap of {coin_name} is ${market_cap:,}"
|
| 15 |
+
MARKET_CAP_FAILURE_MESSAGE = "Failed to retrieve market cap. Please enter a valid coin name."
|
| 16 |
+
API_ERROR_MESSAGE = "I can't seem to access the API at the moment."
|
| 17 |
+
|
src/agents/crypto_data/tools.py
ADDED
|
@@ -0,0 +1,436 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
import requests
|
| 3 |
+
import json
|
| 4 |
+
from sklearn.feature_extraction.text import TfidfVectorizer
|
| 5 |
+
from sklearn.metrics.pairwise import cosine_similarity
|
| 6 |
+
from src.agents.crypto_data.config import Config
|
| 7 |
+
from langchain_core.tools import Tool
|
| 8 |
+
from src.agents.metadata import metadata
|
| 9 |
+
|
| 10 |
+
# -----------------------------------------------------------------------------
|
| 11 |
+
# Module: crypto_data tools
|
| 12 |
+
# -----------------------------------------------------------------------------
|
| 13 |
+
# This module provides helper functions and LangChain Tool wrappers to:
|
| 14 |
+
# - Fetch cryptocurrency prices, market data, and fully diluted valuation
|
| 15 |
+
# - Retrieve NFT floor prices
|
| 16 |
+
# - Query DeFi protocol TVL from DefiLlama
|
| 17 |
+
# - Perform fuzzy matching on protocol names
|
| 18 |
+
#
|
| 19 |
+
# Usage:
|
| 20 |
+
# from src.agents.crypto_data.tools import get_tools
|
| 21 |
+
# tools = get_tools()
|
| 22 |
+
# agent = create_react_agent(model=llm, tools=tools)
|
| 23 |
+
# -----------------------------------------------------------------------------
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def get_most_similar(text: str, data: list[str]) -> list[str]:
|
| 27 |
+
"""
|
| 28 |
+
Find the top entries in `data` most semantically similar to `text`.
|
| 29 |
+
Uses TF-IDF vectorization + cosine similarity.
|
| 30 |
+
|
| 31 |
+
Args:
|
| 32 |
+
text: The input string to match against.
|
| 33 |
+
data: A list of candidate strings.
|
| 34 |
+
|
| 35 |
+
Returns:
|
| 36 |
+
A list of candidates where similarity > 0.5 (up to 20 items).
|
| 37 |
+
"""
|
| 38 |
+
vectorizer = TfidfVectorizer()
|
| 39 |
+
sentence_vectors = vectorizer.fit_transform(data)
|
| 40 |
+
text_vector = vectorizer.transform([text])
|
| 41 |
+
|
| 42 |
+
# Compute cosine similarity between input and all candidates
|
| 43 |
+
similarity_scores = cosine_similarity(text_vector, sentence_vectors)
|
| 44 |
+
|
| 45 |
+
# Pick top 20 indices, then filter by threshold
|
| 46 |
+
top_indices = similarity_scores.argsort()[0][-20:]
|
| 47 |
+
top_matches = [data[i] for i in top_indices if similarity_scores[0][i] > 0.5]
|
| 48 |
+
return top_matches
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def get_coingecko_id(text: str, type: str = "coin") -> str | None:
|
| 52 |
+
"""
|
| 53 |
+
Look up the CoinGecko internal ID for a coin or NFT by name.
|
| 54 |
+
|
| 55 |
+
Args:
|
| 56 |
+
text: Human-readable coin or NFT name or slug.
|
| 57 |
+
type: Either 'coin' or 'nft'.
|
| 58 |
+
|
| 59 |
+
Returns:
|
| 60 |
+
The CoinGecko ID string, or None if not found.
|
| 61 |
+
|
| 62 |
+
Raises:
|
| 63 |
+
ValueError if `type` is invalid, or propagates request exceptions.
|
| 64 |
+
"""
|
| 65 |
+
url = f"{Config.COINGECKO_BASE_URL}/search"
|
| 66 |
+
params = {"query": text}
|
| 67 |
+
try:
|
| 68 |
+
response = requests.get(url, params=params)
|
| 69 |
+
response.raise_for_status()
|
| 70 |
+
data = response.json()
|
| 71 |
+
|
| 72 |
+
if type == "coin":
|
| 73 |
+
return data["coins"][0]["id"] if data["coins"] else None
|
| 74 |
+
elif type == "nft":
|
| 75 |
+
return data.get("nfts", [])[0].get("id") if data.get("nfts") else None
|
| 76 |
+
else:
|
| 77 |
+
raise ValueError("Invalid type specified")
|
| 78 |
+
|
| 79 |
+
except requests.exceptions.RequestException as e:
|
| 80 |
+
logging.error(f"API request failed: {e}")
|
| 81 |
+
raise
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def get_tradingview_symbol(coingecko_id: str) -> str | None:
|
| 85 |
+
"""
|
| 86 |
+
Convert a CoinGecko coin ID into a TradingView ticker symbol.
|
| 87 |
+
|
| 88 |
+
Args:
|
| 89 |
+
coingecko_id: The CoinGecko coin ID.
|
| 90 |
+
|
| 91 |
+
Returns:
|
| 92 |
+
A string like 'CRYPTO:BTCUSD', or None if symbol is missing.
|
| 93 |
+
"""
|
| 94 |
+
url = f"{Config.COINGECKO_BASE_URL}/coins/{coingecko_id}"
|
| 95 |
+
try:
|
| 96 |
+
response = requests.get(url)
|
| 97 |
+
response.raise_for_status()
|
| 98 |
+
symbol = response.json().get("symbol", "").upper()
|
| 99 |
+
return f"CRYPTO:{symbol}USD" if symbol else None
|
| 100 |
+
except requests.exceptions.RequestException as e:
|
| 101 |
+
logging.error(f"Failed to get TradingView symbol: {e}")
|
| 102 |
+
raise
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def get_price(coin: str) -> float | None:
|
| 106 |
+
"""
|
| 107 |
+
Fetch the current USD price of a cryptocurrency.
|
| 108 |
+
|
| 109 |
+
Args:
|
| 110 |
+
coin: Human-readable coin name (e.g. 'bitcoin').
|
| 111 |
+
|
| 112 |
+
Returns:
|
| 113 |
+
Price in USD as a float, or None if coin not found.
|
| 114 |
+
"""
|
| 115 |
+
coin_id = get_coingecko_id(coin, type="coin")
|
| 116 |
+
if not coin_id:
|
| 117 |
+
return None
|
| 118 |
+
|
| 119 |
+
url = f"{Config.COINGECKO_BASE_URL}/simple/price"
|
| 120 |
+
params = {"ids": coin_id, "vs_currencies": "USD"}
|
| 121 |
+
try:
|
| 122 |
+
response = requests.get(url, params=params)
|
| 123 |
+
response.raise_for_status()
|
| 124 |
+
return response.json()[coin_id]["usd"]
|
| 125 |
+
except requests.exceptions.RequestException as e:
|
| 126 |
+
logging.error(f"Failed to retrieve price: {e}")
|
| 127 |
+
raise
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def get_floor_price(nft: str) -> float | None:
|
| 131 |
+
"""
|
| 132 |
+
Retrieve the floor price in USD for an NFT collection.
|
| 133 |
+
|
| 134 |
+
Args:
|
| 135 |
+
nft: The NFT collection name or slug.
|
| 136 |
+
|
| 137 |
+
Returns:
|
| 138 |
+
Floor price in USD, or None if not found.
|
| 139 |
+
"""
|
| 140 |
+
nft_id = get_coingecko_id(nft, type="nft")
|
| 141 |
+
if not nft_id:
|
| 142 |
+
return None
|
| 143 |
+
|
| 144 |
+
url = f"{Config.COINGECKO_BASE_URL}/nfts/{nft_id}"
|
| 145 |
+
try:
|
| 146 |
+
response = requests.get(url)
|
| 147 |
+
response.raise_for_status()
|
| 148 |
+
return response.json()["floor_price"]["usd"]
|
| 149 |
+
except requests.exceptions.RequestException as e:
|
| 150 |
+
logging.error(f"Failed to retrieve floor price: {e}")
|
| 151 |
+
raise
|
| 152 |
+
|
| 153 |
+
|
| 154 |
+
def get_fdv(coin: str) -> float | None:
|
| 155 |
+
"""
|
| 156 |
+
Get a coin's Fully Diluted Valuation (FDV) in USD from CoinGecko.
|
| 157 |
+
|
| 158 |
+
Args:
|
| 159 |
+
coin: Coin name or slug.
|
| 160 |
+
|
| 161 |
+
Returns:
|
| 162 |
+
FDV in USD, or None if not available.
|
| 163 |
+
"""
|
| 164 |
+
coin_id = get_coingecko_id(coin, type="coin")
|
| 165 |
+
if not coin_id:
|
| 166 |
+
return None
|
| 167 |
+
|
| 168 |
+
url = f"{Config.COINGECKO_BASE_URL}/coins/{coin_id}"
|
| 169 |
+
try:
|
| 170 |
+
data = requests.get(url).json()
|
| 171 |
+
return data.get("market_data", {}).get("fully_diluted_valuation", {}).get("usd")
|
| 172 |
+
except requests.exceptions.RequestException as e:
|
| 173 |
+
logging.error(f"Failed to retrieve FDV: {e}")
|
| 174 |
+
raise
|
| 175 |
+
|
| 176 |
+
|
| 177 |
+
def get_market_cap(coin: str) -> float | None:
|
| 178 |
+
"""
|
| 179 |
+
Fetch current market capitalization for a coin via CoinGecko.
|
| 180 |
+
|
| 181 |
+
Args:
|
| 182 |
+
coin: The coin name or slug.
|
| 183 |
+
|
| 184 |
+
Returns:
|
| 185 |
+
Market cap in USD, or None if not found.
|
| 186 |
+
"""
|
| 187 |
+
coin_id = get_coingecko_id(coin, type="coin")
|
| 188 |
+
if not coin_id:
|
| 189 |
+
return None
|
| 190 |
+
|
| 191 |
+
url = f"{Config.COINGECKO_BASE_URL}/coins/markets"
|
| 192 |
+
params = {"ids": coin_id, "vs_currency": "USD"}
|
| 193 |
+
try:
|
| 194 |
+
response = requests.get(url, params=params)
|
| 195 |
+
response.raise_for_status()
|
| 196 |
+
return response.json()[0]["market_cap"]
|
| 197 |
+
except requests.exceptions.RequestException as e:
|
| 198 |
+
logging.error(f"Failed to retrieve market cap: {e}")
|
| 199 |
+
raise
|
| 200 |
+
|
| 201 |
+
|
| 202 |
+
def get_protocols_list() -> tuple[list[str], list[str], list[str]]:
|
| 203 |
+
"""
|
| 204 |
+
Pull the full list of DeFi protocols from DefiLlama.
|
| 205 |
+
|
| 206 |
+
Returns:
|
| 207 |
+
- slugs: List of protocol slugs (for TVL lookup)
|
| 208 |
+
- names: Human-readable names
|
| 209 |
+
- gecko_ids: CoinGecko IDs for integration
|
| 210 |
+
"""
|
| 211 |
+
url = f"{Config.DEFILLAMA_BASE_URL}/protocols"
|
| 212 |
+
try:
|
| 213 |
+
data = requests.get(url).json()
|
| 214 |
+
slugs = [item["slug"] for item in data]
|
| 215 |
+
names = [item["name"] for item in data]
|
| 216 |
+
gecko_ids = [item["gecko_id"] for item in data]
|
| 217 |
+
return slugs, names, gecko_ids
|
| 218 |
+
except requests.exceptions.RequestException as e:
|
| 219 |
+
logging.error(f"Failed to retrieve protocols list: {e}")
|
| 220 |
+
raise
|
| 221 |
+
|
| 222 |
+
|
| 223 |
+
def get_tvl_value(protocol_id: str) -> float:
|
| 224 |
+
"""
|
| 225 |
+
Query DefiLlama for a single protocol's TVL.
|
| 226 |
+
|
| 227 |
+
Args:
|
| 228 |
+
protocol_id: The slug identifier for the protocol.
|
| 229 |
+
|
| 230 |
+
Returns:
|
| 231 |
+
TVL value (could be a dict or number depending on API).
|
| 232 |
+
"""
|
| 233 |
+
url = f"{Config.DEFILLAMA_BASE_URL}/chains"
|
| 234 |
+
try:
|
| 235 |
+
print(f"URL: {url}")
|
| 236 |
+
response = requests.get(url)
|
| 237 |
+
print(f"Response: {response.json()}")
|
| 238 |
+
response.raise_for_status()
|
| 239 |
+
chains = response.json()
|
| 240 |
+
chain = next((c for c in chains if c["name"].lower() == protocol_id.lower()), None)
|
| 241 |
+
print(f"Chain: {chain}")
|
| 242 |
+
return chain["tvl"]
|
| 243 |
+
except requests.exceptions.RequestException as e:
|
| 244 |
+
logging.error(f"Failed to retrieve protocol TVL: {e}")
|
| 245 |
+
raise
|
| 246 |
+
|
| 247 |
+
|
| 248 |
+
def get_protocol_tvl(protocol_name: str) -> dict[str, float] | None:
|
| 249 |
+
"""
|
| 250 |
+
Get a protocol's TVL by name, using fuzzy matching if needed.
|
| 251 |
+
|
| 252 |
+
1. Try exact match via CoinGecko ID → DefiLlama slug
|
| 253 |
+
2. If no exact match, find closest names via TF-IDF
|
| 254 |
+
3. Return the highest TVL among matches
|
| 255 |
+
"""
|
| 256 |
+
slugs, names, gecko_ids = get_protocols_list()
|
| 257 |
+
tag = get_coingecko_id(protocol_name)
|
| 258 |
+
protocol_id = None
|
| 259 |
+
|
| 260 |
+
if tag:
|
| 261 |
+
# map gecko_id to DefiLlama slug
|
| 262 |
+
protocol_id = next((s for s, g in zip(slugs, gecko_ids) if g == tag), None)
|
| 263 |
+
if protocol_id:
|
| 264 |
+
return {tag: get_tvl_value(protocol_id)}
|
| 265 |
+
|
| 266 |
+
# fallback: fuzzy text matching on protocol names
|
| 267 |
+
matches = get_most_similar(protocol_name, names)
|
| 268 |
+
if not matches:
|
| 269 |
+
return None
|
| 270 |
+
|
| 271 |
+
# fetch TVL for each matched name, pick the highest
|
| 272 |
+
results = []
|
| 273 |
+
for name in matches:
|
| 274 |
+
pid = next(s for s, n in zip(slugs, names) if n == name)
|
| 275 |
+
tvl = get_tvl_value(pid)
|
| 276 |
+
results.append({pid: tvl})
|
| 277 |
+
|
| 278 |
+
return max(results, key=lambda d: next(iter(d.values())))
|
| 279 |
+
|
| 280 |
+
|
| 281 |
+
# -----------------------------------------------------------------------------
|
| 282 |
+
# Tool wrappers: these catch errors, format responses, and produce strings
|
| 283 |
+
# -----------------------------------------------------------------------------
|
| 284 |
+
|
| 285 |
+
def _append_coin_metadata_suffix(text: str, coin_name: str) -> str:
|
| 286 |
+
"""
|
| 287 |
+
Append a structured metadata sentinel to a human-friendly text response.
|
| 288 |
+
Includes CoinGecko coinId and uppercased symbol when available.
|
| 289 |
+
"""
|
| 290 |
+
try:
|
| 291 |
+
coin_id = get_coingecko_id(coin_name, type="coin")
|
| 292 |
+
if not coin_id:
|
| 293 |
+
return text
|
| 294 |
+
tv_symbol = get_tradingview_symbol(coin_id)
|
| 295 |
+
symbol = None
|
| 296 |
+
if tv_symbol and tv_symbol.startswith("CRYPTO:") and tv_symbol.endswith("USD"):
|
| 297 |
+
symbol = tv_symbol[len("CRYPTO:"):-3]
|
| 298 |
+
meta = {"coinId": coin_id}
|
| 299 |
+
if symbol:
|
| 300 |
+
meta["symbol"] = symbol
|
| 301 |
+
response = f"{text} ||META: {json.dumps(meta)}||"
|
| 302 |
+
return response
|
| 303 |
+
except Exception:
|
| 304 |
+
# Never break user-visible responses due to metadata failures
|
| 305 |
+
return text
|
| 306 |
+
|
| 307 |
+
|
| 308 |
+
def get_coin_price_tool(coin_name: str) -> dict:
|
| 309 |
+
"""
|
| 310 |
+
LangChain Tool: Return a user-friendly string with the coin's USD price.
|
| 311 |
+
"""
|
| 312 |
+
try:
|
| 313 |
+
price = get_price(coin_name)
|
| 314 |
+
if price is None:
|
| 315 |
+
return Config.PRICE_FAILURE_MESSAGE
|
| 316 |
+
text = Config.PRICE_SUCCESS_MESSAGE.format(coin_name=coin_name, price=price)
|
| 317 |
+
# compute metadata out-of-band
|
| 318 |
+
meta = {}
|
| 319 |
+
try:
|
| 320 |
+
coin_id = get_coingecko_id(coin_name, type="coin")
|
| 321 |
+
if coin_id:
|
| 322 |
+
tv_symbol = get_tradingview_symbol(coin_id)
|
| 323 |
+
meta["coinId"] = tv_symbol # return the symbol as the coinId
|
| 324 |
+
metadata.set_crypto_data_agent(meta)
|
| 325 |
+
except Exception:
|
| 326 |
+
pass
|
| 327 |
+
|
| 328 |
+
return {"text": text, "metadata": meta}
|
| 329 |
+
except requests.exceptions.RequestException:
|
| 330 |
+
return {"text": Config.API_ERROR_MESSAGE, "metadata": {}}
|
| 331 |
+
|
| 332 |
+
|
| 333 |
+
def get_nft_floor_price_tool(nft_name: str) -> str:
|
| 334 |
+
"""
|
| 335 |
+
LangChain Tool: Return a user-friendly string with the NFT floor price.
|
| 336 |
+
"""
|
| 337 |
+
try:
|
| 338 |
+
floor_price = get_floor_price(nft_name)
|
| 339 |
+
if floor_price is None:
|
| 340 |
+
return Config.FLOOR_PRICE_FAILURE_MESSAGE
|
| 341 |
+
return Config.FLOOR_PRICE_SUCCESS_MESSAGE.format(nft_name=nft_name, floor_price=floor_price)
|
| 342 |
+
except requests.exceptions.RequestException:
|
| 343 |
+
return Config.API_ERROR_MESSAGE
|
| 344 |
+
|
| 345 |
+
|
| 346 |
+
def get_protocol_total_value_locked_tool(protocol_name: str) -> str:
|
| 347 |
+
"""
|
| 348 |
+
LangChain Tool: Return formatted TVL information for a DeFi protocol.
|
| 349 |
+
"""
|
| 350 |
+
try:
|
| 351 |
+
tvl = get_protocol_tvl(protocol_name)
|
| 352 |
+
if tvl is None:
|
| 353 |
+
return Config.TVL_FAILURE_MESSAGE
|
| 354 |
+
tag, tvl_value = next(iter(tvl.items()))
|
| 355 |
+
return Config.TVL_SUCCESS_MESSAGE.format(protocol_name=protocol_name, tvl=tvl_value)
|
| 356 |
+
except requests.exceptions.RequestException:
|
| 357 |
+
return Config.API_ERROR_MESSAGE
|
| 358 |
+
|
| 359 |
+
|
| 360 |
+
def get_fully_diluted_valuation_tool(coin_name: str) -> str:
|
| 361 |
+
"""
|
| 362 |
+
LangChain Tool: Return a formatted string with the coin's FDV.
|
| 363 |
+
"""
|
| 364 |
+
try:
|
| 365 |
+
fdv = get_fdv(coin_name)
|
| 366 |
+
if fdv is None:
|
| 367 |
+
return Config.FDV_FAILURE_MESSAGE
|
| 368 |
+
text = Config.FDV_SUCCESS_MESSAGE.format(coin_name=coin_name, fdv=fdv)
|
| 369 |
+
return _append_coin_metadata_suffix(text, coin_name)
|
| 370 |
+
except requests.exceptions.RequestException:
|
| 371 |
+
return Config.API_ERROR_MESSAGE
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
def get_coin_market_cap_tool(coin_name: str) -> str:
|
| 375 |
+
"""
|
| 376 |
+
LangChain Tool: Return a formatted string with the coin's market cap.
|
| 377 |
+
"""
|
| 378 |
+
try:
|
| 379 |
+
market_cap = get_market_cap(coin_name)
|
| 380 |
+
if market_cap is None:
|
| 381 |
+
return Config.MARKET_CAP_FAILURE_MESSAGE
|
| 382 |
+
text = Config.MARKET_CAP_SUCCESS_MESSAGE.format(coin_name=coin_name, market_cap=market_cap)
|
| 383 |
+
return _append_coin_metadata_suffix(text, coin_name)
|
| 384 |
+
except requests.exceptions.RequestException:
|
| 385 |
+
return Config.API_ERROR_MESSAGE
|
| 386 |
+
|
| 387 |
+
|
| 388 |
+
def get_tools() -> list[Tool]:
|
| 389 |
+
"""
|
| 390 |
+
Build and return the list of LangChain Tools for use in an agent.
|
| 391 |
+
|
| 392 |
+
Each Tool wraps one of the user-facing helper functions above.
|
| 393 |
+
"""
|
| 394 |
+
|
| 395 |
+
return [
|
| 396 |
+
Tool(
|
| 397 |
+
name="get_coin_price",
|
| 398 |
+
func=get_coin_price_tool,
|
| 399 |
+
description=(
|
| 400 |
+
"Use this to get the current USD price of a cryptocurrency. "
|
| 401 |
+
"Input should be the coin name (e.g. 'bitcoin')."
|
| 402 |
+
),
|
| 403 |
+
),
|
| 404 |
+
Tool(
|
| 405 |
+
name="get_nft_floor_price",
|
| 406 |
+
func=get_nft_floor_price_tool,
|
| 407 |
+
description=(
|
| 408 |
+
"Fetch the floor price of an NFT collection in USD. "
|
| 409 |
+
"Input should be the NFT name or slug."
|
| 410 |
+
),
|
| 411 |
+
),
|
| 412 |
+
Tool(
|
| 413 |
+
name="get_protocol_tvl",
|
| 414 |
+
func=get_protocol_total_value_locked_tool,
|
| 415 |
+
description=(
|
| 416 |
+
"Returns the Total Value Locked (TVL) of a DeFi protocol. "
|
| 417 |
+
"Input is the protocol name."
|
| 418 |
+
),
|
| 419 |
+
),
|
| 420 |
+
Tool(
|
| 421 |
+
name="get_fully_diluted_valuation",
|
| 422 |
+
func=get_fully_diluted_valuation_tool,
|
| 423 |
+
description=(
|
| 424 |
+
"Get a coin's fully diluted valuation in USD. "
|
| 425 |
+
"Input the coin's name."
|
| 426 |
+
),
|
| 427 |
+
),
|
| 428 |
+
Tool(
|
| 429 |
+
name="get_market_cap",
|
| 430 |
+
func=get_coin_market_cap_tool,
|
| 431 |
+
description=(
|
| 432 |
+
"Retrieve the market capitalization of a coin in USD. "
|
| 433 |
+
"Input is the coin's name."
|
| 434 |
+
),
|
| 435 |
+
),
|
| 436 |
+
]
|
src/agents/database/agent.py
ADDED
|
@@ -0,0 +1,109 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from src.agents.database.tools import get_tools
|
| 2 |
+
from langgraph.prebuilt import create_react_agent
|
| 3 |
+
from langchain_core.prompts import PromptTemplate
|
| 4 |
+
from langchain_core.messages import BaseMessage, SystemMessage, HumanMessage
|
| 5 |
+
from langchain_core.language_models import BaseChatModel
|
| 6 |
+
from langchain_core.runnables import Runnable, RunnableConfig
|
| 7 |
+
from datetime import datetime
|
| 8 |
+
from typing import List, Any
|
| 9 |
+
from src.agents.database.tools import call_tool
|
| 10 |
+
|
| 11 |
+
SYSTEM_PROMPT = """
|
| 12 |
+
You are a senior data assistant specializing in ClickHouse. Your role is to help users query a database related to the Avalanche (AVAX) blockchain network using natural language.
|
| 13 |
+
|
| 14 |
+
You should:
|
| 15 |
+
- Interpret user questions.
|
| 16 |
+
- Explore the database using tools (e.g., list_tables, describe_table, sample_table).
|
| 17 |
+
- Strategically decide which tools to call and when.
|
| 18 |
+
- Generate efficient and accurate SQL queries.
|
| 19 |
+
- Summarize and explain the query results in business-friendly language.
|
| 20 |
+
|
| 21 |
+
## 🧠 STRATEGIC THINKING (BEFORE ACTING)
|
| 22 |
+
- Before calling tools, reason step-by-step.
|
| 23 |
+
- Identify what information is missing.
|
| 24 |
+
- Formulate a plan to investigate the schema or validate assumptions.
|
| 25 |
+
- Only execute SQL after validating the database structure.
|
| 26 |
+
|
| 27 |
+
## 🔧 TOOL USAGE RULES
|
| 28 |
+
- Always include a clear and thoughtful `reasoning` parameter for every tool call.
|
| 29 |
+
- Ensure all required parameters are included and accurate.
|
| 30 |
+
- Use tools sparingly and with purpose. Avoid unnecessary calls.
|
| 31 |
+
- Do not repeat tool calls. After receiving a response, do not call the same tool again unless the user asks for more information.
|
| 32 |
+
|
| 33 |
+
## 🗃️ DATABASE CONTEXT
|
| 34 |
+
- All data is related to the Avalanche (AVAX) blockchain.
|
| 35 |
+
- Tables may include smart contracts, transactions, wallet addresses, gas fees, staking, governance, and on-chain activity.
|
| 36 |
+
|
| 37 |
+
## 📊 OUTPUT FORMAT
|
| 38 |
+
- Always respond in string format.
|
| 39 |
+
- Use bullet points when presenting structured data.
|
| 40 |
+
- If the query involves multiple steps or complex logic, break it down for the user.
|
| 41 |
+
- Assume the user is a **business analyst** or **data scientist** who does **not know SQL**.
|
| 42 |
+
|
| 43 |
+
## 🎯 GOAL
|
| 44 |
+
Transform a vague user request into:
|
| 45 |
+
1. A strategic plan.
|
| 46 |
+
2. The right tool calls to understand the database.
|
| 47 |
+
3. An optimized SQL query.
|
| 48 |
+
4. A clear, insightful explanation of the results.
|
| 49 |
+
|
| 50 |
+
Today’s date is {datetime.now().strftime('%Y-%m-%d')}.
|
| 51 |
+
""".strip()
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class DatabaseAgent(Runnable):
|
| 55 |
+
"""Agent for handling database queries."""
|
| 56 |
+
|
| 57 |
+
def __init__(self, llm, max_iterations: int = 10):
|
| 58 |
+
self.max_iterations = max_iterations
|
| 59 |
+
self.llm = llm.bind_tools(get_tools())
|
| 60 |
+
self.name = "database_agent"
|
| 61 |
+
|
| 62 |
+
def create_history(self) -> List[BaseMessage]:
|
| 63 |
+
"""Create a history of messages for the agent."""
|
| 64 |
+
return [
|
| 65 |
+
SystemMessage(content=SYSTEM_PROMPT),
|
| 66 |
+
]
|
| 67 |
+
|
| 68 |
+
def invoke(self, input: Any, config: RunnableConfig = None) -> str:
|
| 69 |
+
try:
|
| 70 |
+
"""Process the input, which can be a dict or a list of messages."""
|
| 71 |
+
# 🔧 Suporte tanto para dict com chave "messages" quanto para lista direta
|
| 72 |
+
if isinstance(input, dict) and "messages" in input:
|
| 73 |
+
user_messages = input["messages"]
|
| 74 |
+
elif isinstance(input, list):
|
| 75 |
+
user_messages = input
|
| 76 |
+
else:
|
| 77 |
+
raise ValueError("Invalid input format. Expected dict with 'messages' or a list of messages.")
|
| 78 |
+
|
| 79 |
+
n_iterations = 0
|
| 80 |
+
messages = self.create_history() + user_messages # ✅ inclui o system prompt no histórico
|
| 81 |
+
|
| 82 |
+
# initial_response = self.llm.invoke(messages)
|
| 83 |
+
# return {
|
| 84 |
+
# "messages": initial_response,
|
| 85 |
+
# "agent": self.name
|
| 86 |
+
# }
|
| 87 |
+
|
| 88 |
+
while n_iterations < self.max_iterations:
|
| 89 |
+
response = self.llm.invoke(messages)
|
| 90 |
+
print("response", response)
|
| 91 |
+
messages.append(response)
|
| 92 |
+
if not response.tool_calls:
|
| 93 |
+
final_response = {
|
| 94 |
+
"messages": messages,
|
| 95 |
+
"agent": self.name
|
| 96 |
+
}
|
| 97 |
+
print("DEBUG: final_response", final_response)
|
| 98 |
+
return final_response
|
| 99 |
+
for tool_call in response.tool_calls:
|
| 100 |
+
tool_result = call_tool(tool_call)
|
| 101 |
+
messages.append(tool_result)
|
| 102 |
+
n_iterations += 1
|
| 103 |
+
|
| 104 |
+
return response.content
|
| 105 |
+
except Exception as e:
|
| 106 |
+
print(f"Error in DatabaseAgent: {e}")
|
| 107 |
+
return "Sorry, an error occurred while processing your request."
|
| 108 |
+
|
| 109 |
+
|
src/agents/database/client.py
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import clickhouse_connect
|
| 2 |
+
from src.agents.database.config import Config
|
| 3 |
+
|
| 4 |
+
_client = None
|
| 5 |
+
|
| 6 |
+
def _create_client():
|
| 7 |
+
return clickhouse_connect.get_client(
|
| 8 |
+
host=Config.CLICKHOUSE_HOST,
|
| 9 |
+
port=Config.CLICKHOUSE_PORT,
|
| 10 |
+
username=Config.CLICKHOUSE_USER,
|
| 11 |
+
password=Config.CLICKHOUSE_PASSWORD,
|
| 12 |
+
database=Config.CLICKHOUSE_DATABASE
|
| 13 |
+
)
|
| 14 |
+
|
| 15 |
+
def get_client():
|
| 16 |
+
global _client
|
| 17 |
+
if _client is None:
|
| 18 |
+
_client = _create_client()
|
| 19 |
+
return _client
|
| 20 |
+
|
| 21 |
+
def try_get_client():
|
| 22 |
+
try:
|
| 23 |
+
return get_client()
|
| 24 |
+
except Exception:
|
| 25 |
+
return None
|
| 26 |
+
|
| 27 |
+
def is_database_available() -> bool:
|
| 28 |
+
try:
|
| 29 |
+
client = get_client()
|
| 30 |
+
client.query("SELECT 1")
|
| 31 |
+
return True
|
| 32 |
+
except Exception:
|
| 33 |
+
return False
|
| 34 |
+
|
| 35 |
+
def execute_query(query: str):
|
| 36 |
+
return get_client().query(query)
|
src/agents/database/config.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
|
| 3 |
+
class Config:
|
| 4 |
+
# Use clickhouse+http:// for HTTP interface or clickhouse+native:// for native interface
|
| 5 |
+
CLICKHOUSE_URI = "clickhouse+http://default:@localhost:8123/default"
|
| 6 |
+
|
| 7 |
+
CLICKHOUSE_HOST = 'localhost'
|
| 8 |
+
CLICKHOUSE_PORT = 8123 # HTTP port
|
| 9 |
+
CLICKHOUSE_USER = 'default'
|
| 10 |
+
CLICKHOUSE_PASSWORD = ''
|
| 11 |
+
CLICKHOUSE_DATABASE = 'default'
|
| 12 |
+
# Alternative native connection (faster, recommended for production)
|
| 13 |
+
# CLICKHOUSE_URI = "clickhouse+native://default:@localhost:9000/default"
|
| 14 |
+
GLACIER_API_KEY = os.getenv("GLACIER_API_KEY")
|
src/agents/database/tools.py
ADDED
|
@@ -0,0 +1,179 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from langchain_core.messages.tool import ToolCall
|
| 2 |
+
from langchain_core.messages import ToolMessage
|
| 3 |
+
from langchain.tools import Tool, tool
|
| 4 |
+
from src.agents.database.client import get_client
|
| 5 |
+
import requests
|
| 6 |
+
from src.agents.database.config import Config
|
| 7 |
+
|
| 8 |
+
@tool(parse_docstring=True)
|
| 9 |
+
def list_tables(reasoning: str) -> list[str]:
|
| 10 |
+
"""
|
| 11 |
+
List all tables in the database.
|
| 12 |
+
|
| 13 |
+
Args:
|
| 14 |
+
reasoning: A string with the reasoning for the tool call
|
| 15 |
+
|
| 16 |
+
Returns:
|
| 17 |
+
A list of table names
|
| 18 |
+
"""
|
| 19 |
+
print("reasoning:", reasoning)
|
| 20 |
+
client = get_client()
|
| 21 |
+
result = client.query("SHOW TABLES")
|
| 22 |
+
response = result.result_rows
|
| 23 |
+
print("response:", response)
|
| 24 |
+
return response
|
| 25 |
+
|
| 26 |
+
@tool(parse_docstring=True)
|
| 27 |
+
def sample_table(reasoning: str, table_name: str) -> str:
|
| 28 |
+
"""
|
| 29 |
+
Sample the data from a table.
|
| 30 |
+
|
| 31 |
+
Args:
|
| 32 |
+
reasoning: A string with the reasoning for the tool call
|
| 33 |
+
table_name: The name of the table to sample
|
| 34 |
+
|
| 35 |
+
Returns:
|
| 36 |
+
A string with the sampled data
|
| 37 |
+
"""
|
| 38 |
+
print("reasoning:", reasoning)
|
| 39 |
+
client = get_client()
|
| 40 |
+
result = client.query(f"SELECT * FROM {table_name} LIMIT 10")
|
| 41 |
+
response = result.result_rows
|
| 42 |
+
print("response:", response)
|
| 43 |
+
return response
|
| 44 |
+
|
| 45 |
+
@tool(parse_docstring=True)
|
| 46 |
+
def describe_table(reasoning: str, table_name: str) -> str:
|
| 47 |
+
"""
|
| 48 |
+
Describe the schema of a table.
|
| 49 |
+
|
| 50 |
+
Args:
|
| 51 |
+
reasoning: A string with the reasoning for the tool call
|
| 52 |
+
table_name: The name of the table to describe
|
| 53 |
+
|
| 54 |
+
Returns:
|
| 55 |
+
A string with the schema of the table
|
| 56 |
+
"""
|
| 57 |
+
print("reasoning:", reasoning)
|
| 58 |
+
client = get_client()
|
| 59 |
+
result = client.query(f"DESCRIBE TABLE {table_name}")
|
| 60 |
+
response = result.result_rows
|
| 61 |
+
print("response:", response)
|
| 62 |
+
return response
|
| 63 |
+
|
| 64 |
+
@tool(parse_docstring=True)
|
| 65 |
+
def execute_sql(reasoning: str, sql: str) -> str:
|
| 66 |
+
"""
|
| 67 |
+
Execute a SQL query and return the result.
|
| 68 |
+
|
| 69 |
+
Args:
|
| 70 |
+
reasoning: A string with the reasoning for the tool call
|
| 71 |
+
sql: The SQL query to execute
|
| 72 |
+
|
| 73 |
+
Returns:
|
| 74 |
+
A string with the result of the query
|
| 75 |
+
"""
|
| 76 |
+
print("reasoning: ", reasoning)
|
| 77 |
+
client = get_client()
|
| 78 |
+
result = client.query(sql)
|
| 79 |
+
response = result.result_rows
|
| 80 |
+
print("response:", response)
|
| 81 |
+
return response
|
| 82 |
+
|
| 83 |
+
@tool(parse_docstring=True)
|
| 84 |
+
def get_recent_transactions(reasoning: str, blockchainId: str = "c-chain") -> str:
|
| 85 |
+
"""
|
| 86 |
+
Get the recent transactions from the AVAX blockchain database.
|
| 87 |
+
|
| 88 |
+
Args:
|
| 89 |
+
reasoning: A string with the reasoning for the tool call
|
| 90 |
+
blockchainId: The blockchain ID to get the transactions from (default: c-chain)
|
| 91 |
+
|
| 92 |
+
Returns:
|
| 93 |
+
A string with the recent transactions
|
| 94 |
+
"""
|
| 95 |
+
print("reasoning: ", reasoning)
|
| 96 |
+
network = "mainnet"
|
| 97 |
+
url = f"https://glacier-api.avax.network/v1/networks/{network}/blockchains/{blockchainId}/transactions"
|
| 98 |
+
|
| 99 |
+
headers = {"x-glacier-api-key": Config.GLACIER_API_KEY}
|
| 100 |
+
|
| 101 |
+
response = requests.get(url, headers=headers)
|
| 102 |
+
|
| 103 |
+
print("response: ", response.json())
|
| 104 |
+
return str(response.json())
|
| 105 |
+
|
| 106 |
+
@tool(parse_docstring=True)
|
| 107 |
+
def get_recent_active_addresses(reasoning: str, blockchainId: str = "c-chain") -> str:
|
| 108 |
+
"""
|
| 109 |
+
Get the recent active addresses from the AVAX blockchain database.
|
| 110 |
+
|
| 111 |
+
Args:
|
| 112 |
+
reasoning: A string with the reasoning for the tool call
|
| 113 |
+
blockchainId: The blockchain ID to get the active addresses from (default: c-chain)
|
| 114 |
+
|
| 115 |
+
Returns:
|
| 116 |
+
A string with the recent transactions and active addresses
|
| 117 |
+
"""
|
| 118 |
+
print("reasoning: ", reasoning)
|
| 119 |
+
network = "mainnet"
|
| 120 |
+
url = f"https://glacier-api.avax.network/v1/networks/{network}/blockchains/{blockchainId}/transactions"
|
| 121 |
+
headers = {"x-glacier-api-key": Config.GLACIER_API_KEY}
|
| 122 |
+
response = requests.get(url, headers=headers)
|
| 123 |
+
print("response: ", response.json())
|
| 124 |
+
return str(response.json())
|
| 125 |
+
|
| 126 |
+
@tool(parse_docstring=True)
|
| 127 |
+
def get_network_info(reasoning: str) -> str:
|
| 128 |
+
"""
|
| 129 |
+
Gets AVAX mainnet network details such as validator and delegator stats.
|
| 130 |
+
|
| 131 |
+
Args:
|
| 132 |
+
reasoning: A string with the reasoning for the tool call
|
| 133 |
+
|
| 134 |
+
Returns:
|
| 135 |
+
A string with the network details
|
| 136 |
+
"""
|
| 137 |
+
print("reasoning: ", reasoning)
|
| 138 |
+
network = "mainnet"
|
| 139 |
+
url = f"https://glacier-api.avax.network/v1/networks/{network}"
|
| 140 |
+
headers = {"x-glacier-api-key": Config.GLACIER_API_KEY}
|
| 141 |
+
response = requests.get(url, headers=headers)
|
| 142 |
+
print("response: ", response.json())
|
| 143 |
+
return str(response.json())
|
| 144 |
+
|
| 145 |
+
@tool(parse_docstring=True)
|
| 146 |
+
def get_total_active_addresses(reasoning: str, ) -> str:
|
| 147 |
+
"""
|
| 148 |
+
Gets the total active addresses from the AVAX network. Use this tool to get the total number of active addresses on the network.
|
| 149 |
+
|
| 150 |
+
Args:
|
| 151 |
+
reasoning: A string with the reasoning for the tool call
|
| 152 |
+
|
| 153 |
+
Returns:
|
| 154 |
+
A string with the cumulative active addresses
|
| 155 |
+
"""
|
| 156 |
+
print("reasoning: ", reasoning)
|
| 157 |
+
chainId = "total"
|
| 158 |
+
metric = "cumulativeAddresses"
|
| 159 |
+
url = f"https://metrics.avax.network/v2/chains/{chainId}/metrics/{metric}"
|
| 160 |
+
headers = {"x-glacier-api-key": Config.GLACIER_API_KEY}
|
| 161 |
+
response = requests.get(url, headers=headers)
|
| 162 |
+
print("response: ", response.json())
|
| 163 |
+
return str(response.json())
|
| 164 |
+
|
| 165 |
+
def get_tools() -> list[Tool]:
|
| 166 |
+
"""
|
| 167 |
+
Build and return the list of LangChain Tools for use in an agent.
|
| 168 |
+
|
| 169 |
+
Each Tool wraps one of the user-facing helper functions above.
|
| 170 |
+
"""
|
| 171 |
+
|
| 172 |
+
return [list_tables, sample_table, describe_table, execute_sql, get_recent_transactions, get_recent_active_addresses, get_network_info, get_total_active_addresses]
|
| 173 |
+
|
| 174 |
+
def call_tool(tool_call: ToolCall) -> any:
|
| 175 |
+
print(tool_call)
|
| 176 |
+
tools_by_name = {tool.name: tool for tool in get_tools()}
|
| 177 |
+
tool = tools_by_name[tool_call["name"]]
|
| 178 |
+
response = tool.invoke(tool_call["args"])
|
| 179 |
+
return ToolMessage(content=response, tool_call_id=tool_call["id"])
|
src/agents/default/agent.py
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from langgraph.prebuilt import create_react_agent
|
| 2 |
+
|
| 3 |
+
class DefaultAgent():
|
| 4 |
+
"""Agent for handling default queries and data retrieval."""
|
| 5 |
+
|
| 6 |
+
def __init__(self, llm):
|
| 7 |
+
self.llm = llm
|
| 8 |
+
|
| 9 |
+
self.agent = create_react_agent(
|
| 10 |
+
model=llm,
|
| 11 |
+
tools=[],
|
| 12 |
+
name="default_agent"
|
| 13 |
+
)
|
src/agents/metadata.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
class Metadata:
|
| 2 |
+
def __init__(self):
|
| 3 |
+
self.crypto_data_agent = {}
|
| 4 |
+
self.swap_agent = {}
|
| 5 |
+
|
| 6 |
+
def get_crypto_data_agent(self):
|
| 7 |
+
return self.crypto_data_agent
|
| 8 |
+
|
| 9 |
+
def set_crypto_data_agent(self, crypto_data_agent):
|
| 10 |
+
self.crypto_data_agent = crypto_data_agent
|
| 11 |
+
|
| 12 |
+
def get_swap_agent(self):
|
| 13 |
+
return self.swap_agent
|
| 14 |
+
|
| 15 |
+
def set_swap_agent(self, swap_agent):
|
| 16 |
+
self.swap_agent = swap_agent
|
| 17 |
+
|
| 18 |
+
metadata = Metadata()
|
src/agents/supervisor/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
|
src/agents/supervisor/agent.py
ADDED
|
@@ -0,0 +1,306 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from langchain_google_genai import ChatGoogleGenerativeAI, GoogleGenerativeAIEmbeddings
|
| 2 |
+
from langgraph_supervisor import create_supervisor
|
| 3 |
+
from src.agents.config import Config
|
| 4 |
+
from typing import TypedDict, Literal, List, Any
|
| 5 |
+
import re
|
| 6 |
+
import json
|
| 7 |
+
from src.agents.metadata import metadata
|
| 8 |
+
|
| 9 |
+
# Agents
|
| 10 |
+
from src.agents.crypto_data.agent import CryptoDataAgent
|
| 11 |
+
from src.agents.database.agent import DatabaseAgent
|
| 12 |
+
from src.agents.default.agent import DefaultAgent
|
| 13 |
+
from src.agents.swap.agent import SwapAgent
|
| 14 |
+
from src.agents.database.client import is_database_available
|
| 15 |
+
|
| 16 |
+
llm = ChatGoogleGenerativeAI(
|
| 17 |
+
model=Config.GEMINI_MODEL,
|
| 18 |
+
temperature=0.7,
|
| 19 |
+
google_api_key=Config.GEMINI_API_KEY
|
| 20 |
+
)
|
| 21 |
+
|
| 22 |
+
embeddings = GoogleGenerativeAIEmbeddings(
|
| 23 |
+
model=Config.GEMINI_EMBEDDING_MODEL,
|
| 24 |
+
google_api_key=Config.GEMINI_API_KEY
|
| 25 |
+
)
|
| 26 |
+
|
| 27 |
+
class ChatMessage(TypedDict):
|
| 28 |
+
role: Literal["system", "user", "assistant"]
|
| 29 |
+
content: str
|
| 30 |
+
|
| 31 |
+
class Supervisor:
|
| 32 |
+
def __init__(self, llm):
|
| 33 |
+
self.llm = llm
|
| 34 |
+
|
| 35 |
+
cryptoDataAgentClass = CryptoDataAgent(llm)
|
| 36 |
+
cryptoDataAgent = cryptoDataAgentClass.agent
|
| 37 |
+
|
| 38 |
+
agents = [cryptoDataAgent]
|
| 39 |
+
available_agents_text = "- crypto_agent: Handles cryptocurrency-related queries like price checks, market data, NFT floor prices, DeFi protocol TVL, etc.\n"
|
| 40 |
+
|
| 41 |
+
# Conditionally include database agent
|
| 42 |
+
if is_database_available():
|
| 43 |
+
databaseAgent = DatabaseAgent(llm)
|
| 44 |
+
agents.append(databaseAgent)
|
| 45 |
+
available_agents_text += "- database_agent: Handles database queries and data analysis. Can search and analyze data from the database.\n"
|
| 46 |
+
else:
|
| 47 |
+
databaseAgent = None
|
| 48 |
+
|
| 49 |
+
swapAgent = SwapAgent(llm)
|
| 50 |
+
agents.append(swapAgent.agent)
|
| 51 |
+
available_agents_text += "- swap_agent: Handles swap operations on the Avalanche network and any other swap question related.\n"
|
| 52 |
+
|
| 53 |
+
defaultAgent = DefaultAgent(llm)
|
| 54 |
+
agents.append(defaultAgent.agent)
|
| 55 |
+
|
| 56 |
+
# Track known agent names for response extraction
|
| 57 |
+
self.known_agent_names = {"crypto_agent", "database_agent", "swap_agent", "default_agent"}
|
| 58 |
+
|
| 59 |
+
# Prepare database guidance text to avoid backslashes in f-string expressions
|
| 60 |
+
if databaseAgent:
|
| 61 |
+
database_instruction = "When a user asks for data analysis, database queries, or information from the database, delegate to the database_agent."
|
| 62 |
+
database_examples = (
|
| 63 |
+
"Examples of database queries to delegate:\n"
|
| 64 |
+
"- \"What are the avax chains?\"\n"
|
| 65 |
+
"- \"What information is available about the AVAX in the database agent?\"\n"
|
| 66 |
+
"- \"What is the total number of activities addresses in AVAX?\"\n"
|
| 67 |
+
"- \"How many transactions are there in the AVAX network?\"\n"
|
| 68 |
+
)
|
| 69 |
+
else:
|
| 70 |
+
database_instruction = "Do not delegate to a database agent; answer best-effort without DB access or ask the user to start the database."
|
| 71 |
+
database_examples = ""
|
| 72 |
+
|
| 73 |
+
# System prompt to guide the supervisor
|
| 74 |
+
system_prompt = f"""You are a helpful supervisor that routes user queries to the appropriate specialized agents.
|
| 75 |
+
|
| 76 |
+
Available agents:
|
| 77 |
+
{available_agents_text}
|
| 78 |
+
|
| 79 |
+
When a user asks about cryptocurrency prices, market data, NFTs, or DeFi protocols, delegate to the crypto_agent.
|
| 80 |
+
{database_instruction}
|
| 81 |
+
For all other queries, respond directly as a helpful assistant.
|
| 82 |
+
|
| 83 |
+
IMPORTANT: your final response should answer the user's query. Use the agents response to answer the user's query if necessary. Avoid returning control-transfer notes like 'Transferring back to supervisor' — return the substantive answer instead.
|
| 84 |
+
|
| 85 |
+
Examples of crypto queries to delegate:
|
| 86 |
+
- "What is the price of ETH?"
|
| 87 |
+
- "What's the market cap of Bitcoin?"
|
| 88 |
+
- "What's the floor price of Bored Apes?"
|
| 89 |
+
- "What's the TVL of Uniswap?"
|
| 90 |
+
|
| 91 |
+
Examples of swap queries to delegate:
|
| 92 |
+
- I wanna make a swap
|
| 93 |
+
- What are the available tokens for swapping?
|
| 94 |
+
- I want to swap 100 USD for AVAX
|
| 95 |
+
|
| 96 |
+
{database_examples}
|
| 97 |
+
|
| 98 |
+
Examples of general queries to handle directly:
|
| 99 |
+
- "Hello, how are you?"
|
| 100 |
+
- "What's the weather like?"
|
| 101 |
+
- "Tell me a joke"
|
| 102 |
+
"""
|
| 103 |
+
|
| 104 |
+
self.supervisor = create_supervisor(
|
| 105 |
+
agents,
|
| 106 |
+
model=llm,
|
| 107 |
+
prompt=system_prompt,
|
| 108 |
+
output_mode="last_message"
|
| 109 |
+
)
|
| 110 |
+
|
| 111 |
+
self.app = self.supervisor.compile()
|
| 112 |
+
|
| 113 |
+
def _is_handoff_text(self, text: str) -> bool:
|
| 114 |
+
if not text:
|
| 115 |
+
return False
|
| 116 |
+
t = text.strip().lower()
|
| 117 |
+
handoff_keywords = [
|
| 118 |
+
"transferring back",
|
| 119 |
+
"transfer back",
|
| 120 |
+
"returning control",
|
| 121 |
+
"handoff",
|
| 122 |
+
"handing back",
|
| 123 |
+
"delegating back",
|
| 124 |
+
"delegate back",
|
| 125 |
+
"passing back",
|
| 126 |
+
"routing back",
|
| 127 |
+
"route back",
|
| 128 |
+
"back to supervisor",
|
| 129 |
+
"supervisor will handle",
|
| 130 |
+
"sending back to supervisor",
|
| 131 |
+
"give control back",
|
| 132 |
+
"control back to supervisor",
|
| 133 |
+
]
|
| 134 |
+
return any(k in t for k in handoff_keywords)
|
| 135 |
+
|
| 136 |
+
def _sanitize_handoff_phrases(self, text: str) -> str:
|
| 137 |
+
if not text:
|
| 138 |
+
return text
|
| 139 |
+
phrases = [
|
| 140 |
+
"transferring back to supervisor",
|
| 141 |
+
"transfer back to supervisor",
|
| 142 |
+
"returning control to supervisor",
|
| 143 |
+
"handing back to supervisor",
|
| 144 |
+
"delegating back to supervisor",
|
| 145 |
+
"delegate back to supervisor",
|
| 146 |
+
"passing back to supervisor",
|
| 147 |
+
"routing back to supervisor",
|
| 148 |
+
"route back to supervisor",
|
| 149 |
+
"back to supervisor",
|
| 150 |
+
"control back to supervisor",
|
| 151 |
+
"supervisor will handle",
|
| 152 |
+
"sending back to supervisor",
|
| 153 |
+
]
|
| 154 |
+
sanitized = text
|
| 155 |
+
for p in phrases:
|
| 156 |
+
# remove phrase case-insensitively, with optional surrounding punctuation/whitespace
|
| 157 |
+
pattern = re.compile(r"\b" + re.escape(p) + r"\b[\s\.,;:!\)]*", re.IGNORECASE)
|
| 158 |
+
sanitized = pattern.sub(" ", sanitized)
|
| 159 |
+
# Normalize whitespace
|
| 160 |
+
sanitized = re.sub(r"\s+", " ", sanitized).strip()
|
| 161 |
+
return sanitized
|
| 162 |
+
|
| 163 |
+
def _get_text_content(self, message: Any) -> str | None:
|
| 164 |
+
content = getattr(message, "content", None)
|
| 165 |
+
if isinstance(content, str):
|
| 166 |
+
return content
|
| 167 |
+
if isinstance(content, list):
|
| 168 |
+
collected: List[str] = []
|
| 169 |
+
for part in content:
|
| 170 |
+
# Try dict form first
|
| 171 |
+
if isinstance(part, dict):
|
| 172 |
+
text = part.get("text") or part.get("content")
|
| 173 |
+
if isinstance(text, str) and text.strip():
|
| 174 |
+
collected.append(text.strip())
|
| 175 |
+
else:
|
| 176 |
+
# Best-effort: objects with 'text' attribute
|
| 177 |
+
text_attr = getattr(part, "text", None)
|
| 178 |
+
if isinstance(text_attr, str) and text_attr.strip():
|
| 179 |
+
collected.append(text_attr.strip())
|
| 180 |
+
if collected:
|
| 181 |
+
return " ".join(collected)
|
| 182 |
+
return None
|
| 183 |
+
|
| 184 |
+
def _extract_payload(self, text: str) -> tuple[dict, str]:
|
| 185 |
+
# Try JSON payload first
|
| 186 |
+
try:
|
| 187 |
+
obj = json.loads(text)
|
| 188 |
+
if isinstance(obj, dict) and "metadata" in obj and "text" in obj:
|
| 189 |
+
return (obj.get("metadata") or {}), str(obj.get("text") or "")
|
| 190 |
+
except Exception:
|
| 191 |
+
pass
|
| 192 |
+
# Fallback to sentinel
|
| 193 |
+
m = re.search(r"\|\|META:\s*(\{.*?\})\|\|", text)
|
| 194 |
+
if m:
|
| 195 |
+
try:
|
| 196 |
+
meta = json.loads(m.group(1))
|
| 197 |
+
except Exception:
|
| 198 |
+
meta = {}
|
| 199 |
+
cleaned = (text[:m.start()] + text[m.end():]).strip()
|
| 200 |
+
cleaned = re.sub(r"\s+", " ", cleaned)
|
| 201 |
+
return meta, cleaned
|
| 202 |
+
return {}, text
|
| 203 |
+
|
| 204 |
+
def _collect_tool_metadata(self, messages_out) -> dict:
|
| 205 |
+
# Prefer last tool-ish content with metadata
|
| 206 |
+
for m in reversed(messages_out):
|
| 207 |
+
t = self._get_text_content(m) or ""
|
| 208 |
+
meta, _ = self._extract_payload(t)
|
| 209 |
+
if meta:
|
| 210 |
+
return meta
|
| 211 |
+
# Optional: artifact support if your langchain version has it
|
| 212 |
+
art = getattr(m, "artifact", None)
|
| 213 |
+
if isinstance(art, dict) and art:
|
| 214 |
+
return art
|
| 215 |
+
return {}
|
| 216 |
+
|
| 217 |
+
def invoke(self, messages: List[ChatMessage]) -> dict:
|
| 218 |
+
from langchain_core.messages import HumanMessage, SystemMessage, AIMessage
|
| 219 |
+
|
| 220 |
+
langchain_messages = []
|
| 221 |
+
for msg in messages:
|
| 222 |
+
if msg.get("role") == "user":
|
| 223 |
+
langchain_messages.append(HumanMessage(content=msg.get("content", "")))
|
| 224 |
+
elif msg.get("role") == "system":
|
| 225 |
+
langchain_messages.append(SystemMessage(content=msg.get("content", "")))
|
| 226 |
+
elif msg.get("role") == "assistant":
|
| 227 |
+
langchain_messages.append(AIMessage(content=msg.get("content", "")))
|
| 228 |
+
|
| 229 |
+
try:
|
| 230 |
+
response = self.app.invoke({"messages": langchain_messages})
|
| 231 |
+
print("DEBUG: response", response)
|
| 232 |
+
except Exception as e:
|
| 233 |
+
print(f"Error in Supervisor: {e}")
|
| 234 |
+
return {
|
| 235 |
+
"messages": [],
|
| 236 |
+
"agent": "supervisor",
|
| 237 |
+
"response": "Sorry, an error occurred while processing your request."
|
| 238 |
+
}
|
| 239 |
+
|
| 240 |
+
messages_out = response.get("messages", []) if isinstance(response, dict) else []
|
| 241 |
+
|
| 242 |
+
# Prefer the last specialized agent message over any router/supervisor meta message
|
| 243 |
+
final_response = None
|
| 244 |
+
final_agent = "supervisor"
|
| 245 |
+
|
| 246 |
+
def choose_content_from_message(m) -> tuple[str | None, str | None]:
|
| 247 |
+
print("m: ", m)
|
| 248 |
+
agent_name = getattr(m, "name", None)
|
| 249 |
+
content_text = self._get_text_content(m)
|
| 250 |
+
if not content_text:
|
| 251 |
+
return None, None
|
| 252 |
+
sanitized = self._sanitize_handoff_phrases(content_text)
|
| 253 |
+
if sanitized and sanitized.strip() and not self._is_handoff_text(sanitized):
|
| 254 |
+
# Prefer sanitized content
|
| 255 |
+
return sanitized, agent_name
|
| 256 |
+
return None, None
|
| 257 |
+
|
| 258 |
+
|
| 259 |
+
# 1) Try to find the last message from a known specialized agent that is not a handoff/route-back note
|
| 260 |
+
for m in reversed(messages_out):
|
| 261 |
+
agent_name = getattr(m, "name", None)
|
| 262 |
+
if agent_name in self.known_agent_names:
|
| 263 |
+
content, agent = choose_content_from_message(m)
|
| 264 |
+
if content:
|
| 265 |
+
final_response = content
|
| 266 |
+
final_agent = agent or agent_name
|
| 267 |
+
print(f'agent: {agent_name} content: {content}')
|
| 268 |
+
break
|
| 269 |
+
|
| 270 |
+
# 2) Fallback: any last message with content that is not a handoff note
|
| 271 |
+
if final_response is None:
|
| 272 |
+
for m in reversed(messages_out):
|
| 273 |
+
content, agent = choose_content_from_message(m)
|
| 274 |
+
if content:
|
| 275 |
+
final_response = content
|
| 276 |
+
if agent:
|
| 277 |
+
final_agent = agent
|
| 278 |
+
break
|
| 279 |
+
|
| 280 |
+
# 3) Last resort: use dict-level fields if present
|
| 281 |
+
if final_response is None:
|
| 282 |
+
if isinstance(response, dict):
|
| 283 |
+
final_response = response.get("response") or "No response available"
|
| 284 |
+
final_agent = response.get("agent", final_agent)
|
| 285 |
+
else:
|
| 286 |
+
final_response = "No response available"
|
| 287 |
+
|
| 288 |
+
cleaned_response = final_response or "Sorry, no meaningful response was returned."
|
| 289 |
+
meta = {}
|
| 290 |
+
if final_agent == "swap_agent":
|
| 291 |
+
meta = {'src': 'AVAX', 'dst': 'USDC', 'amount': '100'}
|
| 292 |
+
elif final_agent == "crypto_agent":
|
| 293 |
+
meta = metadata.get_crypto_data_agent() or {}
|
| 294 |
+
else:
|
| 295 |
+
meta = {}
|
| 296 |
+
print("meta: ", meta)
|
| 297 |
+
print("cleaned_response: ", cleaned_response)
|
| 298 |
+
|
| 299 |
+
print("final_agent: ", final_agent)
|
| 300 |
+
|
| 301 |
+
return {
|
| 302 |
+
"messages": messages_out,
|
| 303 |
+
"agent": final_agent,
|
| 304 |
+
"response": cleaned_response or "Sorry, no meaningful response was returned.",
|
| 305 |
+
"metadata": meta,
|
| 306 |
+
}
|
src/agents/swap/agent.py
ADDED
|
@@ -0,0 +1,16 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from src.agents.swap.tools import get_tools
|
| 3 |
+
from langgraph.prebuilt import create_react_agent
|
| 4 |
+
|
| 5 |
+
logger = logging.getLogger(__name__)
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class SwapAgent:
|
| 9 |
+
"""Agent for handling swap operations and any other swap related questions"""
|
| 10 |
+
def __init__(self, llm):
|
| 11 |
+
self.llm = llm
|
| 12 |
+
self.agent = create_react_agent(
|
| 13 |
+
model=llm,
|
| 14 |
+
tools=get_tools(),
|
| 15 |
+
name="swap_agent"
|
| 16 |
+
)
|
src/agents/swap/config.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
class SwapConfig:
|
| 2 |
+
"""Configuration and simple interface for supported swap tokens.
|
| 3 |
+
|
| 4 |
+
This provides a canonical set of allowed token symbols and helpers to
|
| 5 |
+
normalize and validate user input before executing swaps.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
# Canonical symbols for Avalanche swaps (expand as needed)
|
| 9 |
+
SUPPORTED_TOKENS = {
|
| 10 |
+
"AVAX", # Native token
|
| 11 |
+
"WAVAX", # Wrapped AVAX
|
| 12 |
+
"USDC",
|
| 13 |
+
"USDT",
|
| 14 |
+
"DAI",
|
| 15 |
+
"WBTC",
|
| 16 |
+
"WETH",
|
| 17 |
+
}
|
| 18 |
+
|
| 19 |
+
@classmethod
|
| 20 |
+
def normalize_symbol(cls, symbol: str) -> str:
|
| 21 |
+
"""Return canonical uppercase symbol without surrounding whitespace."""
|
| 22 |
+
return (symbol or "").strip().upper()
|
| 23 |
+
|
| 24 |
+
@classmethod
|
| 25 |
+
def is_supported(cls, symbol: str) -> bool:
|
| 26 |
+
"""Check if a token symbol is supported (case-insensitive)."""
|
| 27 |
+
return cls.normalize_symbol(symbol) in cls.SUPPORTED_TOKENS
|
| 28 |
+
|
| 29 |
+
@classmethod
|
| 30 |
+
def validate_or_raise(cls, symbol: str) -> str:
|
| 31 |
+
"""Validate token symbol and return its canonical form, or raise ValueError."""
|
| 32 |
+
canonical = cls.normalize_symbol(symbol)
|
| 33 |
+
if canonical not in cls.SUPPORTED_TOKENS:
|
| 34 |
+
raise ValueError(
|
| 35 |
+
f"Unsupported token '{symbol}'. Supported tokens: {sorted(cls.SUPPORTED_TOKENS)}"
|
| 36 |
+
)
|
| 37 |
+
return canonical
|
| 38 |
+
|
| 39 |
+
@classmethod
|
| 40 |
+
def list_supported(cls):
|
| 41 |
+
"""Return a sorted list of supported token symbols."""
|
| 42 |
+
return sorted(cls.SUPPORTED_TOKENS)
|
src/agents/swap/tools.py
ADDED
|
@@ -0,0 +1,49 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from langchain_core.tools import tool
|
| 2 |
+
from src.agents.metadata import metadata
|
| 3 |
+
from src.agents.swap.config import SwapConfig
|
| 4 |
+
|
| 5 |
+
@tool
|
| 6 |
+
def swap_avax(amount: float, from_token: str, to_token: str):
|
| 7 |
+
"""
|
| 8 |
+
Swap AVAX for a given amount of tokens
|
| 9 |
+
|
| 10 |
+
Args:
|
| 11 |
+
amount: The amount of tokens to swap
|
| 12 |
+
from_token: The token to swap from
|
| 13 |
+
to_token: The token to swap to
|
| 14 |
+
|
| 15 |
+
Returns:
|
| 16 |
+
The amount of tokens received
|
| 17 |
+
"""
|
| 18 |
+
try:
|
| 19 |
+
canonical_from = SwapConfig.validate_or_raise(from_token)
|
| 20 |
+
canonical_to = SwapConfig.validate_or_raise(to_token)
|
| 21 |
+
except ValueError as e:
|
| 22 |
+
return str(e)
|
| 23 |
+
|
| 24 |
+
print(f"Swapping {amount} {canonical_from} for {canonical_to}")
|
| 25 |
+
meta = {
|
| 26 |
+
"from_token": canonical_from,
|
| 27 |
+
"to_token": canonical_to,
|
| 28 |
+
"amount": amount
|
| 29 |
+
}
|
| 30 |
+
metadata.set_swap_agent(meta)
|
| 31 |
+
return f'Swapped {amount} {canonical_from} for {canonical_to}'
|
| 32 |
+
|
| 33 |
+
@tool
|
| 34 |
+
def get_avaialble_tokens():
|
| 35 |
+
"""
|
| 36 |
+
Get the available tokens for swapping
|
| 37 |
+
"""
|
| 38 |
+
return SwapConfig.list_supported()
|
| 39 |
+
|
| 40 |
+
@tool
|
| 41 |
+
def default_response():
|
| 42 |
+
"""
|
| 43 |
+
Normal response when the user asks for a swap
|
| 44 |
+
"""
|
| 45 |
+
return f'What would you like to swap? The available agents are{SwapConfig.list_supported}'
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def get_tools():
|
| 49 |
+
return [swap_avax, get_avaialble_tokens, default_response]
|
src/app.py
ADDED
|
@@ -0,0 +1,197 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
logging.basicConfig(
|
| 3 |
+
level=logging.DEBUG,
|
| 4 |
+
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
|
| 5 |
+
handlers=[logging.StreamHandler()]
|
| 6 |
+
)
|
| 7 |
+
logging.info("Test log from app.py startup")
|
| 8 |
+
from fastapi import FastAPI, HTTPException, Request
|
| 9 |
+
from fastapi.middleware.cors import CORSMiddleware
|
| 10 |
+
from pydantic import BaseModel
|
| 11 |
+
from typing import List
|
| 12 |
+
import re
|
| 13 |
+
|
| 14 |
+
from src.agents.config import Config
|
| 15 |
+
from src.agents.supervisor.agent import Supervisor
|
| 16 |
+
from src.models.chatMessage import ChatMessage
|
| 17 |
+
from src.routes.chat_manager_routes import router as chat_manager_router
|
| 18 |
+
from src.service.chat_manager import chat_manager_instance
|
| 19 |
+
from src.agents.crypto_data.tools import get_coingecko_id, get_tradingview_symbol
|
| 20 |
+
|
| 21 |
+
# Initialize FastAPI app
|
| 22 |
+
app = FastAPI(title="Zico Agent API", version="1.0")
|
| 23 |
+
|
| 24 |
+
# Enable CORS for local/frontend dev
|
| 25 |
+
app.add_middleware(
|
| 26 |
+
CORSMiddleware,
|
| 27 |
+
allow_origins=["*"],
|
| 28 |
+
allow_credentials=True,
|
| 29 |
+
allow_methods=["*"],
|
| 30 |
+
allow_headers=["*"],
|
| 31 |
+
)
|
| 32 |
+
|
| 33 |
+
# Instantiate Supervisor agent (singleton LLM)
|
| 34 |
+
supervisor = Supervisor(Config.get_llm())
|
| 35 |
+
|
| 36 |
+
class ChatRequest(BaseModel):
|
| 37 |
+
message: ChatMessage
|
| 38 |
+
chain_id: str = "default"
|
| 39 |
+
wallet_address: str = "default"
|
| 40 |
+
conversation_id: str = "default"
|
| 41 |
+
user_id: str = "anonymous"
|
| 42 |
+
|
| 43 |
+
# Lightweight in-memory agent config for frontend integrations
|
| 44 |
+
AVAILABLE_AGENTS = [
|
| 45 |
+
{"name": "default", "human_readable_name": "Default General Purpose", "description": "General chat and meta-queries about agents."},
|
| 46 |
+
{"name": "crypto data", "human_readable_name": "Crypto Data Fetcher", "description": "Real-time cryptocurrency prices, market cap, FDV, TVL."},
|
| 47 |
+
{"name": "token swap", "human_readable_name": "Token Swap Agent", "description": "Swap tokens using supported DEX APIs."},
|
| 48 |
+
{"name": "realtime search", "human_readable_name": "Real-Time Search", "description": "Search the web for recent information."},
|
| 49 |
+
{"name": "dexscreener", "human_readable_name": "DexScreener Analyst", "description": "Fetches and analyzes DEX trading data."},
|
| 50 |
+
{"name": "rugcheck", "human_readable_name": "Token Safety Analyzer", "description": "Analyzes token safety and trends (Solana)."},
|
| 51 |
+
{"name": "imagen", "human_readable_name": "Image Generator", "description": "Generate images from text prompts."},
|
| 52 |
+
{"name": "rag", "human_readable_name": "Document Assistant", "description": "Answer questions about uploaded documents."},
|
| 53 |
+
{"name": "tweet sizzler", "human_readable_name": "Tweet / X-Post Generator", "description": "Generate engaging tweets."},
|
| 54 |
+
{"name": "dca", "human_readable_name": "DCA Strategy Manager", "description": "Plan and manage DCA strategies."},
|
| 55 |
+
{"name": "base", "human_readable_name": "Base Transaction Manager", "description": "Handle transactions on Base network."},
|
| 56 |
+
{"name": "mor rewards", "human_readable_name": "MOR Rewards Tracker", "description": "Track MOR rewards and balances."},
|
| 57 |
+
{"name": "mor claims", "human_readable_name": "MOR Claims Agent", "description": "Claim MOR tokens."},
|
| 58 |
+
]
|
| 59 |
+
|
| 60 |
+
# Default to a small, reasonable subset
|
| 61 |
+
SELECTED_AGENTS = [agent["name"] for agent in AVAILABLE_AGENTS[:6]]
|
| 62 |
+
|
| 63 |
+
# Commands exposed to the ChatInput autocomplete
|
| 64 |
+
AGENT_COMMANDS = [
|
| 65 |
+
{"command": "morpheus", "name": "Default General Purpose", "description": "General assistant for simple queries and meta-questions."},
|
| 66 |
+
{"command": "crypto", "name": "Crypto Data Fetcher", "description": "Get prices, market cap, FDV, TVL and more."},
|
| 67 |
+
{"command": "document", "name": "Document Assistant", "description": "Ask questions about uploaded documents."},
|
| 68 |
+
{"command": "tweet", "name": "Tweet / X-Post Generator", "description": "Create engaging tweets about crypto and web3."},
|
| 69 |
+
{"command": "search", "name": "Real-Time Search", "description": "Search the web for recent events or updates."},
|
| 70 |
+
{"command": "dexscreener", "name": "DexScreener Analyst", "description": "Analyze DEX trading data on supported chains."},
|
| 71 |
+
{"command": "rugcheck", "name": "Token Safety Analyzer", "description": "Check token safety and view trending tokens."},
|
| 72 |
+
{"command": "dca", "name": "DCA Strategy Manager", "description": "Plan a dollar-cost averaging strategy."},
|
| 73 |
+
{"command": "base", "name": "Base Transaction Manager", "description": "Send tokens and swap on Base."},
|
| 74 |
+
{"command": "rewards", "name": "MOR Rewards Tracker", "description": "Check rewards balance and accrual."},
|
| 75 |
+
]
|
| 76 |
+
|
| 77 |
+
# Agents endpoints expected by the frontend
|
| 78 |
+
@app.get("/agents/available")
|
| 79 |
+
def get_available_agents():
|
| 80 |
+
return {
|
| 81 |
+
"selected_agents": SELECTED_AGENTS,
|
| 82 |
+
"available_agents": AVAILABLE_AGENTS,
|
| 83 |
+
}
|
| 84 |
+
|
| 85 |
+
@app.post("/agents/selected")
|
| 86 |
+
async def set_selected_agents(request: Request):
|
| 87 |
+
global SELECTED_AGENTS
|
| 88 |
+
data = await request.json()
|
| 89 |
+
agents = data.get("agents", [])
|
| 90 |
+
# Validate provided names against available agents
|
| 91 |
+
available_names = {a["name"] for a in AVAILABLE_AGENTS}
|
| 92 |
+
valid_agents = [a for a in agents if a in available_names]
|
| 93 |
+
if not valid_agents:
|
| 94 |
+
# Keep previous selection if nothing valid provided
|
| 95 |
+
return {"status": "no_change", "agents": SELECTED_AGENTS}
|
| 96 |
+
# Update selection
|
| 97 |
+
SELECTED_AGENTS = valid_agents[:6]
|
| 98 |
+
return {"status": "success", "agents": SELECTED_AGENTS}
|
| 99 |
+
|
| 100 |
+
@app.get("/agents/commands")
|
| 101 |
+
def get_agent_commands():
|
| 102 |
+
return {"commands": AGENT_COMMANDS}
|
| 103 |
+
|
| 104 |
+
# Map agent runtime names to high-level types for storage/analytics
|
| 105 |
+
def _map_agent_type(agent_name: str) -> str:
|
| 106 |
+
mapping = {
|
| 107 |
+
"crypto_agent": "crypto data",
|
| 108 |
+
"default_agent": "default",
|
| 109 |
+
"database_agent": "analysis",
|
| 110 |
+
"swap_agent": "token swap",
|
| 111 |
+
"supervisor": "supervisor",
|
| 112 |
+
}
|
| 113 |
+
return mapping.get(agent_name, "supervisor")
|
| 114 |
+
|
| 115 |
+
@app.get("/health")
|
| 116 |
+
def health_check():
|
| 117 |
+
return {"status": "ok"}
|
| 118 |
+
|
| 119 |
+
@app.get("/chat/messages")
|
| 120 |
+
def get_messages(request: Request):
|
| 121 |
+
params = request.query_params
|
| 122 |
+
conversation_id = params.get("conversation_id", "default")
|
| 123 |
+
user_id = params.get("user_id", "anonymous")
|
| 124 |
+
return {"messages": chat_manager_instance.get_messages(conversation_id, user_id)}
|
| 125 |
+
|
| 126 |
+
@app.get("/chat/conversations")
|
| 127 |
+
def get_conversations(request: Request):
|
| 128 |
+
params = request.query_params
|
| 129 |
+
user_id = params.get("user_id", "anonymous")
|
| 130 |
+
return {"conversation_ids": chat_manager_instance.get_all_conversation_ids(user_id)}
|
| 131 |
+
|
| 132 |
+
@app.post("/chat")
|
| 133 |
+
def chat(request: ChatRequest):
|
| 134 |
+
print("request: ", request)
|
| 135 |
+
try:
|
| 136 |
+
# Add the user message to the conversation
|
| 137 |
+
chat_manager_instance.add_message(
|
| 138 |
+
message=request.message.dict(),
|
| 139 |
+
conversation_id=request.conversation_id,
|
| 140 |
+
user_id=request.user_id
|
| 141 |
+
)
|
| 142 |
+
|
| 143 |
+
# Get all messages from the conversation to pass to the agent
|
| 144 |
+
conversation_messages = chat_manager_instance.get_messages(
|
| 145 |
+
conversation_id=request.conversation_id,
|
| 146 |
+
user_id=request.user_id
|
| 147 |
+
)
|
| 148 |
+
|
| 149 |
+
# Invoke the supervisor agent with the conversation
|
| 150 |
+
result = supervisor.invoke(conversation_messages)
|
| 151 |
+
|
| 152 |
+
# Add the agent's response to the conversation
|
| 153 |
+
if result and isinstance(result, dict):
|
| 154 |
+
print("result: ", result)
|
| 155 |
+
agent_name = result.get("agent", "supervisor")
|
| 156 |
+
print("agent_name: ", agent_name)
|
| 157 |
+
agent_name = _map_agent_type(agent_name)
|
| 158 |
+
|
| 159 |
+
# Build response metadata and enrich with coin info for crypto price queries
|
| 160 |
+
response_metadata = {"supervisor_result": result}
|
| 161 |
+
# Prefer supervisor-provided metadata
|
| 162 |
+
if isinstance(result, dict) and result.get("metadata"):
|
| 163 |
+
response_metadata.update(result.get("metadata") or {})
|
| 164 |
+
print("response_metadata: ", response_metadata)
|
| 165 |
+
|
| 166 |
+
# Create a ChatMessage from the supervisor response
|
| 167 |
+
response_message = ChatMessage(
|
| 168 |
+
role="assistant",
|
| 169 |
+
content=result.get("response", "No response available"),
|
| 170 |
+
agent_name=agent_name,
|
| 171 |
+
agent_type=_map_agent_type(agent_name),
|
| 172 |
+
metadata=result.get("metadata", {}),
|
| 173 |
+
conversation_id=request.conversation_id,
|
| 174 |
+
user_id=request.user_id,
|
| 175 |
+
requires_action=True if agent_name == "token swap" else False,
|
| 176 |
+
action_type="swap" if agent_name == "token swap" else None
|
| 177 |
+
)
|
| 178 |
+
|
| 179 |
+
# Add the response message to the conversation
|
| 180 |
+
chat_manager_instance.add_message(
|
| 181 |
+
message=response_message.dict(),
|
| 182 |
+
conversation_id=request.conversation_id,
|
| 183 |
+
user_id=request.user_id
|
| 184 |
+
)
|
| 185 |
+
|
| 186 |
+
# Return only the clean response
|
| 187 |
+
return {
|
| 188 |
+
"response": result.get("response", "No response available"),
|
| 189 |
+
"agentName": agent_name
|
| 190 |
+
}
|
| 191 |
+
|
| 192 |
+
return {"response": "No response available", "agent": "supervisor"}
|
| 193 |
+
except Exception as e:
|
| 194 |
+
raise HTTPException(status_code=500, detail=str(e))
|
| 195 |
+
|
| 196 |
+
# Include chat manager router
|
| 197 |
+
app.include_router(chat_manager_router)
|
src/models/chatMessage.py
ADDED
|
@@ -0,0 +1,152 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import TypedDict, Literal, List, Optional, Dict, Any
|
| 2 |
+
from datetime import datetime
|
| 3 |
+
from pydantic import BaseModel, Field
|
| 4 |
+
from enum import Enum
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class MessageRole(str, Enum):
|
| 8 |
+
"""Enum for message roles"""
|
| 9 |
+
SYSTEM = "system"
|
| 10 |
+
USER = "user"
|
| 11 |
+
ASSISTANT = "assistant"
|
| 12 |
+
AGENT = "agent"
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class AgentType(str, Enum):
|
| 16 |
+
"""Enum for different agent types"""
|
| 17 |
+
SUPERVISOR = "supervisor"
|
| 18 |
+
CRYPTO_DATA = "crypto_data"
|
| 19 |
+
GENERAL = "general"
|
| 20 |
+
RESEARCH = "research"
|
| 21 |
+
ANALYSIS = "analysis"
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class MessageStatus(str, Enum):
|
| 25 |
+
"""Enum for message processing status"""
|
| 26 |
+
PENDING = "pending"
|
| 27 |
+
PROCESSING = "processing"
|
| 28 |
+
COMPLETED = "completed"
|
| 29 |
+
FAILED = "failed"
|
| 30 |
+
CANCELLED = "cancelled"
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
class ChatMessage(BaseModel):
|
| 34 |
+
"""Enhanced chat message model for multi-agent conversations"""
|
| 35 |
+
|
| 36 |
+
# Core message fields
|
| 37 |
+
role: MessageRole = Field(..., description="Role of the message sender")
|
| 38 |
+
content: str = Field(..., description="Message content")
|
| 39 |
+
|
| 40 |
+
# Agent-specific fields
|
| 41 |
+
agent_name: Optional[str] = Field(None, description="Name of the agent that processed this message")
|
| 42 |
+
agent_type: Optional[AgentType] = Field(None, description="Type of agent that processed this message")
|
| 43 |
+
|
| 44 |
+
requires_action: bool = Field(default=False, description="Whether this message requires followup")
|
| 45 |
+
action_type: Optional[str] = Field(None, description="Type of action required")
|
| 46 |
+
|
| 47 |
+
# Metadata and context
|
| 48 |
+
metadata: Dict[str, Any] = Field(default_factory=dict, description="Additional metadata")
|
| 49 |
+
timestamp: datetime = Field(default_factory=datetime.utcnow, description="Message timestamp")
|
| 50 |
+
message_id: Optional[str] = Field(None, description="Unique message identifier")
|
| 51 |
+
|
| 52 |
+
# Processing status
|
| 53 |
+
status: MessageStatus = Field(default=MessageStatus.COMPLETED, description="Message processing status")
|
| 54 |
+
error_message: Optional[str] = Field(None, description="Error message if processing failed")
|
| 55 |
+
|
| 56 |
+
# Conversation context
|
| 57 |
+
conversation_id: Optional[str] = Field(None, description="Conversation identifier")
|
| 58 |
+
user_id: Optional[str] = Field(None, description="User identifier")
|
| 59 |
+
|
| 60 |
+
# Tool calls and responses
|
| 61 |
+
tool_calls: Optional[List[Dict[str, Any]]] = Field(None, description="Tool calls made by the agent")
|
| 62 |
+
tool_results: Optional[List[Dict[str, Any]]] = Field(None, description="Results from tool executions")
|
| 63 |
+
|
| 64 |
+
# Multi-turn conversation support
|
| 65 |
+
next_agent: Optional[str] = Field(None, description="Next agent to handle the conversation")
|
| 66 |
+
requires_followup: bool = Field(default=False, description="Whether this message requires followup")
|
| 67 |
+
|
| 68 |
+
class Config:
|
| 69 |
+
use_enum_values = True
|
| 70 |
+
json_encoders = {
|
| 71 |
+
datetime: lambda v: v.isoformat()
|
| 72 |
+
}
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
class ConversationState(BaseModel):
|
| 76 |
+
"""State management for multi-agent conversations"""
|
| 77 |
+
|
| 78 |
+
conversation_id: str = Field(..., description="Unique conversation identifier")
|
| 79 |
+
user_id: str = Field(..., description="User identifier")
|
| 80 |
+
|
| 81 |
+
# Current state
|
| 82 |
+
current_agent: Optional[str] = Field(None, description="Currently active agent")
|
| 83 |
+
last_message_id: Optional[str] = Field(None, description="ID of the last message")
|
| 84 |
+
|
| 85 |
+
# Conversation history
|
| 86 |
+
messages: List[ChatMessage] = Field(default_factory=list, description="Message history")
|
| 87 |
+
|
| 88 |
+
# Context and memory
|
| 89 |
+
context: Dict[str, Any] = Field(default_factory=dict, description="Conversation context")
|
| 90 |
+
memory: Dict[str, Any] = Field(default_factory=dict, description="Persistent memory across turns")
|
| 91 |
+
|
| 92 |
+
# Agent routing history
|
| 93 |
+
agent_history: List[Dict[str, Any]] = Field(default_factory=list, description="History of agent interactions")
|
| 94 |
+
|
| 95 |
+
# Status and metadata
|
| 96 |
+
created_at: datetime = Field(default_factory=datetime.utcnow)
|
| 97 |
+
updated_at: datetime = Field(default_factory=datetime.utcnow)
|
| 98 |
+
is_active: bool = Field(default=True, description="Whether conversation is active")
|
| 99 |
+
|
| 100 |
+
class Config:
|
| 101 |
+
use_enum_values = True
|
| 102 |
+
json_encoders = {
|
| 103 |
+
datetime: lambda v: v.isoformat()
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
class AgentResponse(BaseModel):
|
| 108 |
+
"""Standardized response format for agents"""
|
| 109 |
+
|
| 110 |
+
content: str = Field(..., description="Response content")
|
| 111 |
+
agent_name: str = Field(..., description="Name of the responding agent")
|
| 112 |
+
agent_type: AgentType = Field(..., description="Type of the responding agent")
|
| 113 |
+
|
| 114 |
+
# Metadata
|
| 115 |
+
metadata: Dict[str, Any] = Field(default_factory=dict, description="Response metadata")
|
| 116 |
+
timestamp: datetime = Field(default_factory=datetime.utcnow)
|
| 117 |
+
|
| 118 |
+
# Tool information
|
| 119 |
+
tools_used: List[str] = Field(default_factory=list, description="Tools used in this response")
|
| 120 |
+
tool_results: Optional[List[Dict[str, Any]]] = Field(None, description="Results from tool executions")
|
| 121 |
+
|
| 122 |
+
# Next steps
|
| 123 |
+
next_agent: Optional[str] = Field(None, description="Next agent to handle the conversation")
|
| 124 |
+
requires_followup: bool = Field(default=False, description="Whether followup is needed")
|
| 125 |
+
|
| 126 |
+
# Status
|
| 127 |
+
success: bool = Field(default=True, description="Whether the response was successful")
|
| 128 |
+
error_message: Optional[str] = Field(None, description="Error message if failed")
|
| 129 |
+
|
| 130 |
+
class Config:
|
| 131 |
+
use_enum_values = True
|
| 132 |
+
json_encoders = {
|
| 133 |
+
datetime: lambda v: v.isoformat()
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
# TypedDict for backward compatibility
|
| 138 |
+
class ChatMessageDict(TypedDict):
|
| 139 |
+
role: str
|
| 140 |
+
content: str
|
| 141 |
+
agent_name: Optional[str]
|
| 142 |
+
metadata: Dict[str, Any]
|
| 143 |
+
timestamp: str
|
| 144 |
+
message_id: Optional[str]
|
| 145 |
+
status: str
|
| 146 |
+
conversation_id: Optional[str]
|
| 147 |
+
user_id: Optional[str]
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
|
src/routes/chat_manager_routes.py
ADDED
|
@@ -0,0 +1,90 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from fastapi import APIRouter, Query, Body
|
| 3 |
+
from src.service.chat_manager import chat_manager_instance
|
| 4 |
+
from typing import Optional
|
| 5 |
+
from pydantic import BaseModel
|
| 6 |
+
from typing import Dict
|
| 7 |
+
|
| 8 |
+
logger = logging.getLogger(__name__)
|
| 9 |
+
|
| 10 |
+
router = APIRouter(prefix="/chat", tags=["chat"])
|
| 11 |
+
|
| 12 |
+
class UserIdRequest(BaseModel):
|
| 13 |
+
user_id: str
|
| 14 |
+
|
| 15 |
+
@router.get("/messages")
|
| 16 |
+
async def get_messages(conversation_id: str = Query(default="default"), user_id: str = Query(default="anonymous")):
|
| 17 |
+
"""Get all chat messages for a conversation"""
|
| 18 |
+
logger.info(f"Received get_messages request for conversation {conversation_id} from user {user_id}")
|
| 19 |
+
return {"messages": chat_manager_instance.get_messages(conversation_id, user_id)}
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
@router.get("/clear")
|
| 23 |
+
async def clear_messages(conversation_id: str = Query(default="default"), user_id: str = Query(default="anonymous")):
|
| 24 |
+
"""Clear chat message history for a conversation"""
|
| 25 |
+
logger.info(f"Clearing message history for conversation {conversation_id} for user {user_id}")
|
| 26 |
+
chat_manager_instance.clear_messages(conversation_id, user_id)
|
| 27 |
+
return {"response": "successfully cleared message history"}
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
@router.get("/conversations")
|
| 31 |
+
async def get_conversations(
|
| 32 |
+
user_id_query: str = Query(default=None, alias="user_id"),
|
| 33 |
+
user_id_str: Optional[str] = Body(default=None)
|
| 34 |
+
):
|
| 35 |
+
"""Get all conversation IDs for a specific user"""
|
| 36 |
+
user_id = user_id_str if user_id_str else user_id_query
|
| 37 |
+
|
| 38 |
+
if not user_id:
|
| 39 |
+
user_id = "anonymous"
|
| 40 |
+
|
| 41 |
+
logger.info(f"Getting all conversation IDs for user {user_id}")
|
| 42 |
+
return {"conversation_ids": chat_manager_instance.get_all_conversation_ids(user_id)}
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
@router.get("/users")
|
| 46 |
+
async def get_users():
|
| 47 |
+
"""Get all user IDs"""
|
| 48 |
+
logger.info("Getting all user IDs")
|
| 49 |
+
return {"user_ids": chat_manager_instance.get_all_user_ids()}
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
@router.post("/conversations")
|
| 53 |
+
async def create_conversation(
|
| 54 |
+
user_id_query: str = Query(default=None, alias="user_id"),
|
| 55 |
+
user_id_body: Optional[UserIdRequest] = None,
|
| 56 |
+
user_id_str: Optional[str] = Body(default=None)
|
| 57 |
+
):
|
| 58 |
+
"""Create a new conversation for a specific user"""
|
| 59 |
+
user_id = None
|
| 60 |
+
if user_id_body:
|
| 61 |
+
user_id = user_id_body.user_id
|
| 62 |
+
elif user_id_str:
|
| 63 |
+
user_id = user_id_str
|
| 64 |
+
else:
|
| 65 |
+
user_id = user_id_query
|
| 66 |
+
|
| 67 |
+
if not user_id:
|
| 68 |
+
user_id = "anonymous"
|
| 69 |
+
|
| 70 |
+
logger.info(f"Creating new conversation for user {user_id}")
|
| 71 |
+
conversation_id = chat_manager_instance.create_conversation(user_id)
|
| 72 |
+
return {"conversation_id": conversation_id}
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
@router.delete("/conversations/{conversation_id}")
|
| 76 |
+
async def delete_conversation(
|
| 77 |
+
conversation_id: str,
|
| 78 |
+
user_id_query: str = Query(default=None, alias="user_id"),
|
| 79 |
+
user_id_body: Dict[str, str] = Body(default=None)
|
| 80 |
+
):
|
| 81 |
+
"""Delete a conversation by ID"""
|
| 82 |
+
# Get user_id from body if it exists, otherwise from query
|
| 83 |
+
user_id = user_id_body.get("user_id") if user_id_body else user_id_query
|
| 84 |
+
|
| 85 |
+
if not user_id:
|
| 86 |
+
user_id = "anonymous"
|
| 87 |
+
|
| 88 |
+
logger.info(f"Deleting conversation {conversation_id} for user {user_id}")
|
| 89 |
+
chat_manager_instance.delete_conversation(conversation_id, user_id)
|
| 90 |
+
return {"response": "successfully deleted conversation"}
|
src/service/chat_manager.py
ADDED
|
@@ -0,0 +1,280 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
import time
|
| 3 |
+
from typing import Dict, List, Optional
|
| 4 |
+
from src.models.chatMessage import ChatMessage, ConversationState, AgentResponse
|
| 5 |
+
from datetime import datetime
|
| 6 |
+
|
| 7 |
+
logger = logging.getLogger(__name__)
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class ChatManager:
|
| 11 |
+
"""
|
| 12 |
+
Manages chat conversations and message history.
|
| 13 |
+
|
| 14 |
+
This class provides functionality to:
|
| 15 |
+
- Create and manage multiple conversations identified by unique IDs
|
| 16 |
+
- Associate conversations with specific users
|
| 17 |
+
- Add/retrieve messages and responses within conversations
|
| 18 |
+
- Clear conversation history
|
| 19 |
+
- Get chat history in different formats
|
| 20 |
+
- Delete conversations
|
| 21 |
+
|
| 22 |
+
Each conversation starts with a default disclaimer message about the experimental nature
|
| 23 |
+
of the chatbot.
|
| 24 |
+
|
| 25 |
+
Attributes:
|
| 26 |
+
user_conversations (Dict[str, Dict[str, ConversationState]]): Dictionary mapping user IDs to their conversations
|
| 27 |
+
default_message (ChatMessage): Default disclaimer message added to new conversations
|
| 28 |
+
|
| 29 |
+
Example:
|
| 30 |
+
>>> chat_manager = ChatManager()
|
| 31 |
+
>>> chat_manager.add_message({"role": "user", "content": "Hello"}, "conv1", "user123")
|
| 32 |
+
>>> messages = chat_manager.get_messages("conv1", "user123")
|
| 33 |
+
"""
|
| 34 |
+
|
| 35 |
+
def __init__(self) -> None:
|
| 36 |
+
self.user_conversations: Dict[str, Dict[str, ConversationState]] = {}
|
| 37 |
+
self.default_message = ChatMessage(
|
| 38 |
+
role="assistant",
|
| 39 |
+
content="""This highly experimental chatbot is not intended for making important decisions. Its
|
| 40 |
+
responses are generated using AI models and may not always be accurate.
|
| 41 |
+
By using this chatbot, you acknowledge that you use it at your own discretion
|
| 42 |
+
and assume all risks associated with its limitations and potential errors.""",
|
| 43 |
+
metadata={},
|
| 44 |
+
)
|
| 45 |
+
|
| 46 |
+
# Backward compatibility - Initialize with default conversation for anonymous user
|
| 47 |
+
self._initialize_user("anonymous")
|
| 48 |
+
self.user_conversations["anonymous"]["default"] = ConversationState(
|
| 49 |
+
conversation_id="default",
|
| 50 |
+
user_id="anonymous",
|
| 51 |
+
messages=[self.default_message],
|
| 52 |
+
context={},
|
| 53 |
+
memory={},
|
| 54 |
+
agent_history=[],
|
| 55 |
+
current_agent=None,
|
| 56 |
+
last_message_id=None,
|
| 57 |
+
created_at=datetime.utcnow(),
|
| 58 |
+
updated_at=datetime.utcnow(),
|
| 59 |
+
is_active=True
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
def _get_conversation_id(self, conversation_id: Optional[str] = None) -> str:
|
| 63 |
+
"""Helper method to get conversation ID, defaulting to 'default' if None provided"""
|
| 64 |
+
return conversation_id or "default"
|
| 65 |
+
|
| 66 |
+
def _get_user_id(self, user_id: Optional[str] = None) -> str:
|
| 67 |
+
"""Helper method to get user ID, defaulting to 'anonymous' if None provided"""
|
| 68 |
+
return user_id or "anonymous"
|
| 69 |
+
|
| 70 |
+
def _initialize_user(self, user_id: str) -> None:
|
| 71 |
+
"""Initialize conversations dictionary for a new user"""
|
| 72 |
+
if user_id not in self.user_conversations:
|
| 73 |
+
self.user_conversations[user_id] = {}
|
| 74 |
+
logger.info(f"Initialized conversations for user {user_id}")
|
| 75 |
+
|
| 76 |
+
def get_messages(self, conversation_id: Optional[str] = None, user_id: Optional[str] = None) -> List[Dict[str, str]]:
|
| 77 |
+
"""
|
| 78 |
+
Get all messages for a specific conversation.
|
| 79 |
+
|
| 80 |
+
Args:
|
| 81 |
+
conversation_id (str, optional): Unique identifier for the conversation. Defaults to "default"
|
| 82 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 83 |
+
|
| 84 |
+
Returns:
|
| 85 |
+
List[Dict[str, str]]: List of messages as dictionaries
|
| 86 |
+
"""
|
| 87 |
+
user_id = self._get_user_id(user_id)
|
| 88 |
+
conversation = self._get_or_create_conversation(self._get_conversation_id(conversation_id), user_id)
|
| 89 |
+
return [msg.dict() for msg in conversation.messages]
|
| 90 |
+
|
| 91 |
+
def add_message(self, message: Dict[str, str], conversation_id: Optional[str] = None, user_id: Optional[str] = None):
|
| 92 |
+
"""
|
| 93 |
+
Add a new message to a conversation.
|
| 94 |
+
|
| 95 |
+
Args:
|
| 96 |
+
message (Dict[str, str]): Message to add
|
| 97 |
+
conversation_id (str, optional): Conversation to add message to. Defaults to "default"
|
| 98 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 99 |
+
"""
|
| 100 |
+
user_id = self._get_user_id(user_id)
|
| 101 |
+
conversation_id = self._get_conversation_id(conversation_id)
|
| 102 |
+
conversation = self._get_or_create_conversation(conversation_id, user_id)
|
| 103 |
+
chat_message = ChatMessage(**message)
|
| 104 |
+
if "timestamp" not in message:
|
| 105 |
+
chat_message.timestamp = datetime.utcnow()
|
| 106 |
+
conversation.messages.append(chat_message)
|
| 107 |
+
logger.info(f"Added message to conversation {conversation_id} for user {user_id}: {chat_message.content}")
|
| 108 |
+
|
| 109 |
+
def add_response(self, response: Dict[str, str], agent_name: str, conversation_id: Optional[str] = None, user_id: Optional[str] = None):
|
| 110 |
+
"""
|
| 111 |
+
Add an agent's response to a conversation.
|
| 112 |
+
|
| 113 |
+
Args:
|
| 114 |
+
response (Dict[str, str]): Response content
|
| 115 |
+
agent_name (str): Name of the responding agent
|
| 116 |
+
conversation_id (str, optional): Conversation to add response to. Defaults to "default"
|
| 117 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 118 |
+
"""
|
| 119 |
+
agent_response = AgentResponse(**response)
|
| 120 |
+
# You may need to implement to_chat_message if not present
|
| 121 |
+
chat_message = ChatMessage(
|
| 122 |
+
role="assistant",
|
| 123 |
+
content=agent_response.content,
|
| 124 |
+
agent_name=agent_response.agent_name,
|
| 125 |
+
agent_type=agent_response.agent_type,
|
| 126 |
+
metadata=agent_response.metadata,
|
| 127 |
+
timestamp=agent_response.timestamp,
|
| 128 |
+
tool_results=agent_response.tool_results,
|
| 129 |
+
next_agent=agent_response.next_agent,
|
| 130 |
+
requires_followup=agent_response.requires_followup,
|
| 131 |
+
status="completed" if agent_response.success else "failed",
|
| 132 |
+
error_message=agent_response.error_message
|
| 133 |
+
)
|
| 134 |
+
self.add_message(chat_message.dict(), self._get_conversation_id(conversation_id), self._get_user_id(user_id))
|
| 135 |
+
logger.info(f"Added response from agent {agent_name} to conversation {conversation_id} for user {user_id}")
|
| 136 |
+
|
| 137 |
+
def clear_messages(self, conversation_id: Optional[str] = None, user_id: Optional[str] = None):
|
| 138 |
+
"""
|
| 139 |
+
Clear all messages in a conversation except the default message.
|
| 140 |
+
|
| 141 |
+
Args:
|
| 142 |
+
conversation_id (str, optional): Conversation to clear. Defaults to "default"
|
| 143 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 144 |
+
"""
|
| 145 |
+
user_id = self._get_user_id(user_id)
|
| 146 |
+
conversation = self._get_or_create_conversation(self._get_conversation_id(conversation_id), user_id)
|
| 147 |
+
conversation.messages = [self.default_message] # Keep the initial message
|
| 148 |
+
logger.info(f"Cleared message history for conversation {conversation_id} for user {user_id}")
|
| 149 |
+
|
| 150 |
+
def get_last_message(self, conversation_id: Optional[str] = None, user_id: Optional[str] = None) -> Dict[str, str]:
|
| 151 |
+
"""
|
| 152 |
+
Get the most recent message from a conversation.
|
| 153 |
+
|
| 154 |
+
Args:
|
| 155 |
+
conversation_id (str, optional): Conversation to get message from. Defaults to "default"
|
| 156 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 157 |
+
|
| 158 |
+
Returns:
|
| 159 |
+
Dict[str, str]: Last message or empty dict if no messages
|
| 160 |
+
"""
|
| 161 |
+
user_id = self._get_user_id(user_id)
|
| 162 |
+
conversation = self._get_or_create_conversation(self._get_conversation_id(conversation_id), user_id)
|
| 163 |
+
return conversation.messages[-1].dict() if conversation.messages else {}
|
| 164 |
+
|
| 165 |
+
def get_chat_history(self, conversation_id: Optional[str] = None, user_id: Optional[str] = None) -> str:
|
| 166 |
+
"""
|
| 167 |
+
Get formatted chat history for a conversation.
|
| 168 |
+
|
| 169 |
+
Args:
|
| 170 |
+
conversation_id (str, optional): Conversation to get history for. Defaults to "default"
|
| 171 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 172 |
+
|
| 173 |
+
Returns:
|
| 174 |
+
str: Formatted chat history as string
|
| 175 |
+
"""
|
| 176 |
+
user_id = self._get_user_id(user_id)
|
| 177 |
+
conversation = self._get_or_create_conversation(self._get_conversation_id(conversation_id), user_id)
|
| 178 |
+
return "\n".join([f"{msg.role}: {msg.content}" for msg in conversation.messages])
|
| 179 |
+
|
| 180 |
+
def get_all_conversation_ids(self, user_id: Optional[str] = None) -> List[str]:
|
| 181 |
+
"""
|
| 182 |
+
Get a list of all conversation IDs for a specific user.
|
| 183 |
+
|
| 184 |
+
Args:
|
| 185 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 186 |
+
|
| 187 |
+
Returns:
|
| 188 |
+
List[str]: List of conversation IDs for the user
|
| 189 |
+
"""
|
| 190 |
+
user_id = self._get_user_id(user_id)
|
| 191 |
+
self._initialize_user(user_id)
|
| 192 |
+
return list(self.user_conversations[user_id].keys())
|
| 193 |
+
|
| 194 |
+
def get_all_user_ids(self) -> List[str]:
|
| 195 |
+
"""
|
| 196 |
+
Get a list of all user IDs.
|
| 197 |
+
|
| 198 |
+
Returns:
|
| 199 |
+
List[str]: List of user IDs
|
| 200 |
+
"""
|
| 201 |
+
return list(self.user_conversations.keys())
|
| 202 |
+
|
| 203 |
+
def delete_conversation(self, conversation_id: Optional[str] = None, user_id: Optional[str] = None):
|
| 204 |
+
"""
|
| 205 |
+
Delete a conversation by ID.
|
| 206 |
+
|
| 207 |
+
Args:
|
| 208 |
+
conversation_id (str, optional): ID of conversation to delete. Defaults to "default"
|
| 209 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 210 |
+
"""
|
| 211 |
+
user_id = self._get_user_id(user_id)
|
| 212 |
+
conversation_id = self._get_conversation_id(conversation_id)
|
| 213 |
+
self._initialize_user(user_id)
|
| 214 |
+
if conversation_id in self.user_conversations[user_id]:
|
| 215 |
+
del self.user_conversations[user_id][conversation_id]
|
| 216 |
+
logger.info(f"Deleted conversation {conversation_id} for user {user_id}")
|
| 217 |
+
|
| 218 |
+
def create_conversation(self, user_id: Optional[str] = None) -> str:
|
| 219 |
+
"""
|
| 220 |
+
Create a new conversation for a user.
|
| 221 |
+
|
| 222 |
+
Args:
|
| 223 |
+
user_id (str, optional): User identifier. Defaults to "anonymous"
|
| 224 |
+
|
| 225 |
+
Returns:
|
| 226 |
+
str: ID of the created conversation
|
| 227 |
+
"""
|
| 228 |
+
user_id = self._get_user_id(user_id)
|
| 229 |
+
self._initialize_user(user_id)
|
| 230 |
+
conversation_id = f"conversation_{len(self.user_conversations[user_id])}"
|
| 231 |
+
while conversation_id in self.user_conversations[user_id]:
|
| 232 |
+
conversation_id = f"conversation_{len(self.user_conversations[user_id])}_{int(time.time())}"
|
| 233 |
+
self.user_conversations[user_id][conversation_id] = ConversationState(
|
| 234 |
+
conversation_id=conversation_id,
|
| 235 |
+
user_id=user_id,
|
| 236 |
+
messages=[self.default_message],
|
| 237 |
+
context={},
|
| 238 |
+
memory={},
|
| 239 |
+
agent_history=[],
|
| 240 |
+
current_agent=None,
|
| 241 |
+
last_message_id=None,
|
| 242 |
+
created_at=datetime.utcnow(),
|
| 243 |
+
updated_at=datetime.utcnow(),
|
| 244 |
+
is_active=True
|
| 245 |
+
)
|
| 246 |
+
logger.info(f"Created new conversation {conversation_id} for user {user_id}")
|
| 247 |
+
return conversation_id
|
| 248 |
+
|
| 249 |
+
def _get_or_create_conversation(self, conversation_id: str, user_id: str) -> ConversationState:
|
| 250 |
+
"""
|
| 251 |
+
Get existing conversation or create new one if not exists.
|
| 252 |
+
|
| 253 |
+
Args:
|
| 254 |
+
conversation_id (str): Conversation ID to get/create
|
| 255 |
+
user_id (str): User identifier
|
| 256 |
+
|
| 257 |
+
Returns:
|
| 258 |
+
ConversationState: Retrieved or created conversation
|
| 259 |
+
"""
|
| 260 |
+
self._initialize_user(user_id)
|
| 261 |
+
if conversation_id not in self.user_conversations[user_id]:
|
| 262 |
+
self.user_conversations[user_id][conversation_id] = ConversationState(
|
| 263 |
+
conversation_id=conversation_id,
|
| 264 |
+
user_id=user_id,
|
| 265 |
+
messages=[self.default_message],
|
| 266 |
+
context={},
|
| 267 |
+
memory={},
|
| 268 |
+
agent_history=[],
|
| 269 |
+
current_agent=None,
|
| 270 |
+
last_message_id=None,
|
| 271 |
+
created_at=datetime.utcnow(),
|
| 272 |
+
updated_at=datetime.utcnow(),
|
| 273 |
+
is_active=True
|
| 274 |
+
)
|
| 275 |
+
logger.info(f"Created new conversation {conversation_id} for user {user_id}")
|
| 276 |
+
return self.user_conversations[user_id][conversation_id]
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
# Create an instance to act as a singleton store
|
| 280 |
+
chat_manager_instance = ChatManager()
|