more fixes

This commit is contained in:
Nicolò Boschi 2025-10-31 17:54:29 +01:00
parent b3259d2907
commit d8d48d6f80
3 changed files with 95 additions and 9 deletions

View file

@ -45,7 +45,7 @@ class LinkInfo(BaseModel):
"""Information about a link to a neighbor.""" """Information about a link to a neighbor."""
to_node_id: str = Field(description="Target node ID") to_node_id: str = Field(description="Target node ID")
link_type: Literal["temporal", "semantic", "entity"] = Field(description="Type of link") link_type: Literal["temporal", "semantic", "entity"] = Field(description="Type of link")
link_weight: float = Field(description="Weight of the link", ge=0.0, le=1.0) link_weight: float = Field(description="Weight of the link (can exceed 1.0 when aggregating multiple connections)", ge=0.0)
entity_id: Optional[str] = Field(default=None, description="Entity ID if link_type is 'entity'") entity_id: Optional[str] = Field(default=None, description="Entity ID if link_type is 'entity'")
new_activation: Optional[float] = Field(default=None, description="Activation that would be passed to neighbor (None for supplementary links)") new_activation: Optional[float] = Field(default=None, description="Activation that would be passed to neighbor (None for supplementary links)")
followed: bool = Field(description="Whether this link was followed (or pruned)") followed: bool = Field(description="Whether this link was followed (or pruned)")

View file

@ -14,7 +14,8 @@ from dotenv import load_dotenv
import os import os
import sys import sys
from pathlib import Path from pathlib import Path
from typing import Optional from typing import Optional, List, Dict, Any
from datetime import datetime
# Add parent directory to path for imports # Add parent directory to path for imports
sys.path.insert(0, str(Path(__file__).parent.parent)) sys.path.insert(0, str(Path(__file__).parent.parent))
@ -38,6 +39,23 @@ class SearchRequest(BaseModel):
thinking_budget: int = 100 thinking_budget: int = 100
top_k: int = 10 top_k: int = 10
mmr_lambda: float = 0.5 mmr_lambda: float = 0.5
trace: bool = False
class MemoryItem(BaseModel):
"""Single memory item for batch put."""
content: str
event_date: Optional[datetime] = None
context: Optional[str] = None
class BatchPutRequest(BaseModel):
"""Request model for batch put endpoint."""
agent_id: str
items: List[MemoryItem]
document_id: Optional[str] = None
document_metadata: Optional[Dict[str, Any]] = None
upsert: bool = False
async def get_graph_data(): async def get_graph_data():
@ -179,6 +197,7 @@ async def get_graph_data():
"total_units": len(units) "total_units": len(units)
} }
memory = TemporalSemanticMemory()
@app.get("/") @app.get("/")
async def index(): async def index():
@ -204,15 +223,13 @@ async def api_search(request: SearchRequest):
"""Run a search and return results with trace.""" """Run a search and return results with trace."""
try: try:
# Initialize memory system # Initialize memory system
memory = TemporalSemanticMemory()
# Run search with tracing # Run search with tracing
results, trace = await memory.search_async( results, trace = await memory.search_async(
agent_id=request.agent_id, agent_id=request.agent_id,
query=request.query, query=request.query,
thinking_budget=request.thinking_budget, thinking_budget=request.thinking_budget,
top_k=request.top_k, top_k=request.top_k,
enable_trace=True, enable_trace=request.trace,
mmr_lambda=request.mmr_lambda mmr_lambda=request.mmr_lambda
) )
@ -258,6 +275,72 @@ async def api_agents():
raise HTTPException(status_code=500, detail=str(e)) raise HTTPException(status_code=500, detail=str(e))
@app.post("/api/memories/batch")
async def api_batch_put(request: BatchPutRequest):
"""
Store multiple memories in batch.
This endpoint calls put_batch_async to efficiently store multiple memory items.
Supports document tracking and upsert operations.
Example request:
{
"agent_id": "user123",
"items": [
{"content": "Alice works at Google", "context": "work"},
{"content": "Bob went hiking yesterday", "event_date": "2024-01-15T10:00:00Z"}
],
"document_id": "conversation_123",
"upsert": false
}
"""
try:
# Validate agent_id - prevent writing to reserved agents
RESERVED_AGENT_IDS = {"locomo"}
if request.agent_id in RESERVED_AGENT_IDS:
raise HTTPException(
status_code=403,
detail=f"Cannot write to reserved agent_id '{request.agent_id}'. Reserved agents: {', '.join(RESERVED_AGENT_IDS)}"
)
# Initialize memory system
# Prepare contents for put_batch_async
contents = []
for item in request.items:
content_dict = {"content": item.content}
if item.event_date:
content_dict["event_date"] = item.event_date
if item.context:
content_dict["context"] = item.context
contents.append(content_dict)
# Call put_batch_async
result = await memory.put_batch_async(
agent_id=request.agent_id,
contents=contents,
document_id=request.document_id,
document_metadata=request.document_metadata,
upsert=request.upsert
)
await memory.close()
return {
"success": True,
"message": f"Successfully stored {len(contents)} memory items",
"agent_id": request.agent_id,
"document_id": request.document_id,
"items_count": len(contents)
}
except Exception as e:
import traceback
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
print(f"Error in /api/memories/batch: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.get("/api/locomo") @app.get("/api/locomo")
async def api_locomo(): async def api_locomo():
"""Get Locomo benchmark results.""" """Get Locomo benchmark results."""
@ -283,9 +366,11 @@ if __name__ == "__main__":
print("=" * 80) print("=" * 80)
print("\nStarting server at http://localhost:8080") print("\nStarting server at http://localhost:8080")
print("\nEndpoints:") print("\nEndpoints:")
print(" GET / - Visualization UI") print(" GET / - Visualization UI")
print(" GET /api/graph - Get graph data") print(" GET /api/graph - Get graph data")
print(" POST /api/search - Run search with trace") print(" POST /api/search - Run search with trace")
print(" POST /api/memories/batch - Store multiple memories in batch")
print(" GET /api/agents - List available agents")
print("\n" + "=" * 80 + "\n") print("\n" + "=" * 80 + "\n")
uvicorn.run("server:app", host="0.0.0.0", port=8080, reload=True) uvicorn.run("server:app", host="0.0.0.0", port=8080, reload=True)

View file

@ -483,7 +483,8 @@ window.runSearchInPane = async function(paneId) {
agent_id: agentId, agent_id: agentId,
thinking_budget: thinkingBudget, thinking_budget: thinkingBudget,
top_k: topK, top_k: topK,
mmr_lambda: mmrLambda mmr_lambda: mmrLambda,
trace: true
}) })
}); });