more fixes
This commit is contained in:
parent
b3259d2907
commit
d8d48d6f80
3 changed files with 95 additions and 9 deletions
|
|
@ -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)")
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
})
|
})
|
||||||
});
|
});
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue