diff --git a/Makefile b/Makefile index 0e575e9..d82bde0 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,4 @@ -.PHONY: dev infra backend frontend up down ingest serve +.PHONY: dev infra backend frontend up down ingest serve lint lint-fix dev: up @@ -41,3 +41,11 @@ migrate: # Run the StackOverflow ingestion pipeline ingest: cd backend && PYTHONPATH=$(PWD)/backend ../.venv/bin/python ingestion/stackoverflow_loader.py + +# Lint the backend with ruff +lint: + cd backend && ../.venv/bin/ruff check . + +# Lint and auto-fix what ruff can fix +lint-fix: + cd backend && ../.venv/bin/ruff check --fix . diff --git a/backend/agents/graph.py b/backend/agents/graph.py index 3c8bf49..23c4736 100644 --- a/backend/agents/graph.py +++ b/backend/agents/graph.py @@ -1,8 +1,10 @@ -from typing import TypedDict, Annotated, Sequence, Any -from langchain_core.messages import BaseMessage import operator import os -from langgraph.graph import StateGraph, END +from collections.abc import Sequence +from typing import Annotated, Any, TypedDict + +from langchain_core.messages import BaseMessage +from langgraph.graph import END, StateGraph REDIS_URL = os.getenv("REDIS_URL", "redis://localhost:6379") @@ -61,11 +63,11 @@ def route_query(state: SynapticState) -> str: # Node imports after SynapticState is defined to avoid circular import failure -from .nodes.triage import triage_node -from .nodes.rag_agent import rag_agent -from .nodes.memory_node import memory_node -from .nodes.orchestrator import orchestrator_node -from .nodes.writer_node import writer_node +from .nodes.memory_node import memory_node # noqa: E402 +from .nodes.orchestrator import orchestrator_node # noqa: E402 +from .nodes.rag_agent import rag_agent # noqa: E402 +from .nodes.triage import triage_node # noqa: E402 +from .nodes.writer_node import writer_node # noqa: E402 # ── Graph definition ────────────────────────────────────────────────────────── diff --git a/backend/agents/nodes/memory_node.py b/backend/agents/nodes/memory_node.py index d0e8d9f..a649bc8 100644 --- a/backend/agents/nodes/memory_node.py +++ b/backend/agents/nodes/memory_node.py @@ -1,7 +1,7 @@ -import asyncio +from langfuse import observe + from agents.graph import SynapticState from memory.manager import MemoryManager -from langfuse import observe memory_manager: MemoryManager = None diff --git a/backend/agents/nodes/orchestrator.py b/backend/agents/nodes/orchestrator.py index e533a9e..64c24ae 100644 --- a/backend/agents/nodes/orchestrator.py +++ b/backend/agents/nodes/orchestrator.py @@ -1,14 +1,17 @@ import asyncio -from typing import Any, Optional -from langchain_core.documents import Document -from langchain_core.messages import SystemMessage, HumanMessage -from llm import utility_llm +from typing import Any + +from langchain_core.messages import HumanMessage, SystemMessage +from langfuse import observe, get_client + from agents.graph import SynapticState +from constants import BEHAVIORAL_GUARDRAILS +from context.engineer import assemble +from llm import utility_llm from memory.manager import MemoryManager + from . import rag_agent as _rag_agent -from .rag_agent import format_memory, extract_citations -from constants import BEHAVIORAL_GUARDRAILS -from langfuse import observe +from .rag_agent import extract_citations ORCHESTRATOR_PROMPT = """ You are a synthesis assistant. The user's query requires both factual knowledge and conversation history. @@ -20,18 +23,7 @@ If conversation history contradicts the documents, prefer the documents but acknowledge the discrepancy. """ + BEHAVIORAL_GUARDRAILS -memory_manager: Optional[MemoryManager] = None - - -def _format_docs(docs: list[Document]) -> str: - return "\n\n".join( - f"[Source: {d.metadata.get('title', 'Unknown')}]\n{d.page_content}" - for d in docs - ) - - -def _format_long_term(long_term: list[dict[str, Any]]) -> str: - return "\n\n".join(m.get("summary", "") for m in long_term if m.get("summary")) +memory_manager: MemoryManager | None = None class Orchestrator: @@ -70,21 +62,35 @@ async def orchestrator_node(state: SynapticState) -> dict[str, Any]: """ assert memory_manager is not None, "memory_manager not initialised" + langfuse = get_client() + memory_result, docs = await asyncio.gather( memory_manager.load(session_id=state["session_id"], query=state["query"]), _rag_agent.retriever.ainvoke(state["query"]), ) - short_term = format_memory(memory_result.get("short_term_memory", [])) - long_term = _format_long_term(memory_result.get("long_term_memory", [])) - docs_context = _format_docs(docs) + bundle = assemble( + "orchestrator_node", + short_term_memory=memory_result.get("short_term_memory", []), + long_term_memory=memory_result.get("long_term_memory", []), + retrieved_chunks=docs, + ) + + langfuse.update_current_span( + metadata={ + "budget_decision": bundle.decision, + "token_count": bundle.token_count, + "budget_exceeded": state.get("budget_exceeded"), + } + ) + callbacks = state.get("callbacks", []) answer = await _orchestrator.merge( query=state["query"], - docs_context=docs_context, - short_term=short_term, - long_term=long_term, + docs_context=bundle.chunks_context, + short_term=bundle.short_term_context, + long_term=bundle.long_term_context, callbacks=callbacks, ) @@ -93,4 +99,8 @@ async def orchestrator_node(state: SynapticState) -> dict[str, Any]: "retrieved_chunks": [{"content": d.page_content, **d.metadata} for d in docs], "citations": extract_citations(docs), "final_answer": answer, + "token_counts": bundle.token_count, + "total_tokens": bundle.total_tokens, + "budget_exceeded": bundle.budget_exceeded, + "decision": bundle.decision, } diff --git a/backend/agents/nodes/rag_agent.py b/backend/agents/nodes/rag_agent.py index 6141342..0dd06c2 100644 --- a/backend/agents/nodes/rag_agent.py +++ b/backend/agents/nodes/rag_agent.py @@ -1,10 +1,13 @@ -from agents.graph import SynapticState -from chain.rag_chain import RagChain +from typing import Any + from langchain_core.documents import Document from langchain_huggingface import HuggingFaceEmbeddings -from typing import Any from langfuse import observe +from agents.graph import SynapticState +from chain.rag_chain import RagChain +from context.engineer import assemble + rag_chain = None retriever = None @@ -17,36 +20,37 @@ def init(embeddings: HuggingFaceEmbeddings) -> None: @observe(name="rag_agent_node") -async def rag_agent(state: SynapticState): - memory_context = format_memory(state["short_term_memory"]) +async def rag_agent(state: SynapticState) -> dict[str, Any]: + + bundle = assemble( + "rag_agent", + short_term_memory=state["short_term_memory"], + long_term_memory=state["long_term_memory"], + ) callbacks = state.get("callbacks", []) result = await rag_chain.ainvoke( - {"question": state["query"], "memory_context": memory_context}, + {"question": state["query"], "memory_context": bundle.memory_context}, config={"callbacks": callbacks}, ) + + token_counts = { + **bundle.token_count, + "retrieved_chunks": result.get("chunks_tokens", 0), + } + return { "final_answer": result["answer"], "retrieved_chunks": result["source_documents"], "citations": extract_citations(result["source_documents"]), "condensed_query": result.get("condensed_query", state["query"]), + "token_counts": token_counts, + "total_tokens": sum(token_counts.values()), + "decision": bundle.decision, + "budget_exceeded": bundle.budget_exceeded + or result.get("chunks_truncated", False), } -def format_memory(short_term_memory: list[dict[str, str]]) -> str: - - memory = "" - for message in short_term_memory: - role: str - - if message["role"] == "ai": - role = "AI" - else: - role = "Human" - memory += f"{role}: {message['content']}\n" - - return memory - - def extract_citations(documents: list[Document]) -> list[dict[str, Any]]: citations = [] for doc in documents: diff --git a/backend/agents/nodes/triage.py b/backend/agents/nodes/triage.py index 94d4788..a28e959 100644 --- a/backend/agents/nodes/triage.py +++ b/backend/agents/nodes/triage.py @@ -1,10 +1,12 @@ -from pydantic import BaseModel -from langchain_core.messages import SystemMessage, HumanMessage -from typing import Any -import re import logging -from agents.graph import SynapticState +import re +from typing import Any + +from langchain_core.messages import HumanMessage, SystemMessage from langfuse import observe +from pydantic import BaseModel + +from agents.graph import SynapticState from llm import utility_llm logger = logging.getLogger(__name__) diff --git a/backend/agents/nodes/writer_node.py b/backend/agents/nodes/writer_node.py index 0a58dcb..f6f065c 100644 --- a/backend/agents/nodes/writer_node.py +++ b/backend/agents/nodes/writer_node.py @@ -1,6 +1,7 @@ +from langfuse import observe + from agents.graph import SynapticState from memory.manager import MemoryManager -from langfuse import observe memory_manager: MemoryManager = None diff --git a/backend/chain/rag_chain.py b/backend/chain/rag_chain.py index c1be601..4dd7e80 100644 --- a/backend/chain/rag_chain.py +++ b/backend/chain/rag_chain.py @@ -1,14 +1,15 @@ from typing import Any -from langchain_huggingface import HuggingFaceEmbeddings -from langchain_core.documents import Document -from langchain_core.prompts import ChatPromptTemplate from langchain_core.output_parsers import StrOutputParser +from langchain_core.prompts import ChatPromptTemplate from langchain_core.runnables import RunnableLambda +from langchain_huggingface import HuggingFaceEmbeddings + +from constants import SYSTEM_PROMPT +from context.engineer import AGENT_BUDGET, fit_chunks from ingestion.stackoverflow_loader import CONN_STR from llm import main_llm, utility_llm from retrieval.chunks_retriever import ChunksRetriever -from constants import SYSTEM_PROMPT _CONDENSATION_SYSTEM_PROMPT = ( "Rewrite the follow-up question as a standalone question by replacing any " @@ -31,12 +32,6 @@ def __init__(self, embeddings: HuggingFaceEmbeddings) -> None: ) self.prompt = self._build_prompt() - def _format_docs(self, docs: list[Document]) -> str: - return "\n\n".join( - f"[Source: {doc.metadata.get('title', 'Unknown')}]\n{doc.page_content}" - for doc in docs - ) - def _build_prompt(self) -> ChatPromptTemplate: return ChatPromptTemplate.from_messages( [ @@ -60,7 +55,7 @@ def _build_condensation_chain(self) -> Any: ) return prompt | utility_llm | StrOutputParser() - def build(self): + def build(self) -> Any: _NO_CONTEXT_REPLY = "I couldn't find relevant information for your question." _condensation_chain = self._build_condensation_chain() _llm_chain = self.prompt | main_llm | StrOutputParser() @@ -100,9 +95,17 @@ async def _retrieve(inputs: dict) -> dict: async def _answer(inputs: dict) -> dict: docs = inputs.get("source_documents", []) if not docs: - return {**inputs, "context": "", "answer": _NO_CONTEXT_REPLY} + return { + **inputs, + "context": "", + "answer": _NO_CONTEXT_REPLY, + "chunks_tokens": 0, + "chunks_truncated": False, + } - context = self._format_docs(docs) + context, chunks_tokens, chunks_truncated = fit_chunks( + docs, AGENT_BUDGET["rag_agent"]["retrieved_chunks"] + ) answer = await _llm_chain.ainvoke( { "question": inputs["question"], @@ -110,7 +113,13 @@ async def _answer(inputs: dict) -> dict: "context": context, } ) - return {**inputs, "answer": answer, "context": context} + return { + **inputs, + "answer": answer, + "context": context, + "chunks_tokens": chunks_tokens, + "chunks_truncated": chunks_truncated, + } return ( RunnableLambda(_condense) diff --git a/backend/context/engineer.py b/backend/context/engineer.py new file mode 100644 index 0000000..b6dcbdf --- /dev/null +++ b/backend/context/engineer.py @@ -0,0 +1,183 @@ +from dataclasses import dataclass +from typing import Any + +import tiktoken +from langchain_core.documents import Document + +from optimiser import Optimiser + +ENCODING = tiktoken.get_encoding("cl100k_base") + +# Unsloth Studio Gemma 4 E4B deployment. +N_CTX = 60_928 + +# Tokens reserved for the model's own generated answer. Enforced via +# llm.py's main_llm(max_tokens=RESERVED_OUTPUT_TOKENS). +RESERVED_OUTPUT_TOKENS = 2_048 + +# Buffer against tiktoken cl100k_base being only an approximation of +# Gemma's real tokenizer (see ARCHITECTURE.md §3). Not enforced at a +# specific call site — it's headroom baked into TOTAL_INPUT_BUDGET below. +SAFETY_MARGIN_TOKENS = 2_048 + +TOTAL_INPUT_BUDGET = N_CTX - RESERVED_OUTPUT_TOKENS - SAFETY_MARGIN_TOKENS + +AGENT_BUDGET = { + "rag_agent": { + "total": TOTAL_INPUT_BUDGET, + "system_prompt": 1_500, + "query_and_prompt_scaffolding": 1_100, + "short_term_memory": 9_000, + "long_term_memory": 6_800, + "retrieved_chunks": 36_000, + "spare": 2_432, + }, + "orchestrator_node": { + "total": TOTAL_INPUT_BUDGET, + "system_prompt": 1_500, + "query_and_prompt_scaffolding": 1_100, + "short_term_memory": 7_900, + "long_term_memory": 11_300, + "retrieved_chunks": 31_600, + "spare": 3_432, + }, +} + + +optimiser = Optimiser() + + +@dataclass +class ContextBundle: + memory_context: str + short_term_context: str + long_term_context: str + chunks_context: str + token_count: dict[str, int] + total_tokens: int + budget_exceeded: bool + decision: list[dict[str, Any]] + + +def count_tokens(text: str) -> int: + return len(ENCODING.encode(text)) + + +def _format_short_term(turns: list[dict[str, Any]]) -> str: + lines: list[str] = [] + for message in turns: + role = "AI" if message["role"] == "ai" else "Human" + lines.append(f"{role}: {message['content']}") + return "\n".join(lines) + + +def _format_long_term(summaries: list[dict[str, Any]]) -> str: + return "\n\n".join(s.get("summary", "") for s in summaries if s.get("summary")) + + +def _format_chunks(docs: list[Document]) -> str: + return "\n\n".join( + f"[Source: {doc.metadata.get('title', 'Unknown')}]\n{doc.page_content}" + for doc in docs + ) + + +def fit_short_term(turns: list[dict[str, Any]], cap: int) -> tuple[str, int, bool]: + """Drop oldest human/assistant pairs first until the formatted text fits cap.""" + truncated = False + while turns: + text = _format_short_term(turns) + tokens = count_tokens(text) + if tokens <= cap: + return text, tokens, truncated + turns = turns[2:] + truncated = True + return "", 0, truncated + + +def fit_long_term(summaries: list[dict[str, Any]], cap: int) -> tuple[str, int, bool]: + """Summaries arrive sorted best-similarity-first (LongTermMemory.load); drop the + lowest-ranked (tail) first.""" + truncated = False + while summaries: + text = _format_long_term(summaries) + tokens = count_tokens(text) + if tokens <= cap: + return text, tokens, truncated + summaries = summaries[:-1] + truncated = True + return "", 0, truncated + + +def fit_chunks(docs: list[Document], cap: int) -> tuple[str, int, bool]: + """Sort the Docs according to the best-relevance-first (ChunksRetriever rerank) using the relevance score; drop the lowest-ranked (tail) first.""" + + docs = sorted(docs, key=lambda d: d.metadata.get("score", 0.0), reverse=True) + truncated = False + original_count = len(docs) + overflow_threshold = original_count * 0.5 + while docs: + text = _format_chunks(docs) + tokens = count_tokens(text) + if tokens <= cap: + return text, tokens, truncated + if len(docs) <= overflow_threshold: + compresssed = optimiser.compress( + [d.page_content for d in docs], target_tokens=cap + ) + return compresssed, count_tokens(compresssed), True + docs = docs[:-1] + truncated = True + return "", 0, truncated + + +def assemble( + agent_name: str, + short_term_memory: list[dict[str, Any]] | None = None, + long_term_memory: list[dict[str, Any]] | None = None, + retrieved_chunks: list[Document] | None = None, +) -> ContextBundle: + budget = AGENT_BUDGET[agent_name] + + short_term_text, st_tokens, st_truncated = fit_short_term( + short_term_memory or [], budget["short_term_memory"] + ) + long_term_text, lt_tokens, lt_truncated = fit_long_term( + long_term_memory or [], budget["long_term_memory"] + ) + chunks_text, chunks_tokens, chunks_truncated = fit_chunks( + retrieved_chunks or [], budget["retrieved_chunks"] + ) + + token_count = { + "short_term_memory": st_tokens, + "long_term_memory": lt_tokens, + "retrieved_chunks": chunks_tokens, + } + + decision: list[dict[str, Any]] = [] + if st_truncated: + decision.append( + {"field": "short_term_memory", "action": "dropped_oldest_pairs"} + ) + if lt_truncated: + decision.append( + {"field": "long_term_memory", "action": "dropped_lowest_relevance"} + ) + if chunks_truncated: + decision.append( + {"field": "retrieved_chunks", "action": "dropped_lowest_relevance"} + ) + + memory_context = "\n\n".join(t for t in (short_term_text, long_term_text) if t) + + return ContextBundle( + memory_context=memory_context, + short_term_context=short_term_text, + long_term_context=long_term_text, + chunks_context=chunks_text, + token_count=token_count, + total_tokens=sum(token_count.values()), + budget_exceeded=bool(decision), + decision=decision, + ) diff --git a/backend/context/optimiser.py b/backend/context/optimiser.py new file mode 100644 index 0000000..fa8b027 --- /dev/null +++ b/backend/context/optimiser.py @@ -0,0 +1,37 @@ +from llmlingua import PromptCompressor +from llm import utility_llm +import logging + +logger = logging.getLogger(__name__) + + +class Optimiser: + + def __init__(self) -> None: + self._compressor = None + + def _get_compressor(self) -> PromptCompressor: + if self._compressor is None: + self._compressor = PromptCompressor( + model_name="microsoft/llmlingua-2-xlm-roberta-large-meetingbank", + use_llmlingua2=True, + device_map="cpu", + ) + + return self._compressor + + def compress(self, texts: list[str], target_tokens: int) -> str: + try: + result = self._get_compressor().compress_prompt( + context=texts, target_token=target_tokens + ) + return str(result["compressed_prompt"]) + except Exception as e: + logger.warning( + msg=f"WARN: Exception Occured. Failed to compress prompt using compressor. Reason: {str(e)}" + ) + prompt = f"Summarise below in ~{target_tokens} tokens:\n\n" + "\n\n".join( + texts + ) + + return utility_llm.invoke(prompt).content diff --git a/backend/ingestion/stackoverflow_data_builder.py b/backend/ingestion/stackoverflow_data_builder.py index 96f4ce0..77d7b14 100644 --- a/backend/ingestion/stackoverflow_data_builder.py +++ b/backend/ingestion/stackoverflow_data_builder.py @@ -1,12 +1,13 @@ import hashlib import logging import warnings +from pathlib import Path from typing import TypedDict + import pandas as pd -from pathlib import Path from bs4 import BeautifulSoup, XMLParsedAsHTMLWarning -from tqdm import tqdm from langchain_core.documents import Document +from tqdm import tqdm warnings.filterwarnings("ignore", category=XMLParsedAsHTMLWarning) diff --git a/backend/ingestion/stackoverflow_loader.py b/backend/ingestion/stackoverflow_loader.py index 6f4aca8..1117e38 100644 --- a/backend/ingestion/stackoverflow_loader.py +++ b/backend/ingestion/stackoverflow_loader.py @@ -1,18 +1,20 @@ -from dotenv import load_dotenv import json import logging import os -import warnings import time +import warnings + import psycopg from bs4 import XMLParsedAsHTMLWarning -from tqdm import tqdm +from dotenv import load_dotenv from langchain_core.documents import Document from langchain_huggingface import HuggingFaceEmbeddings +from tqdm import tqdm + from ingestion.stackoverflow_data_builder import ( - DocumentMetadata, DATA_PATH, EVAL_IDS_PATH, + DocumentMetadata, SODatasetBuilder, ) @@ -300,6 +302,7 @@ def _insert_chunks( texts[batch_slice], vectors[batch_slice], metadatas[batch_slice], + strict=True, ): copy.write_row( ( diff --git a/backend/llm.py b/backend/llm.py index 10a8113..60f012b 100644 --- a/backend/llm.py +++ b/backend/llm.py @@ -1,7 +1,10 @@ import os + from dotenv import load_dotenv -from pydantic import SecretStr from langchain_openai import ChatOpenAI +from pydantic import SecretStr + +from context.engineer import RESERVED_OUTPUT_TOKENS load_dotenv() @@ -16,6 +19,7 @@ api_key=SecretStr(_API_KEY), temperature=1.0, top_p=0.95, + max_tokens=RESERVED_OUTPUT_TOKENS, ) # Used by: Triage, Orchestrator, RagChain (condensation), MemoryManager (summarisation) diff --git a/backend/main.py b/backend/main.py index 210e37f..9cf30e1 100644 --- a/backend/main.py +++ b/backend/main.py @@ -1,46 +1,44 @@ -from contextlib import asynccontextmanager -from fastapi import FastAPI -from fastapi.responses import StreamingResponse -from fastapi.middleware.cors import CORSMiddleware -from llm import utility_llm -from pydantic import BaseModel -from typing import Literal, Optional -import time -import uuid -import json -import os import asyncio +import json import logging import re -from ingestion.stackoverflow_loader import IngestionPipeline +import time +import uuid +from contextlib import asynccontextmanager +from typing import Literal -import redis.asyncio as aioredis import asyncpg +import redis.asyncio as aioredis +from fastapi import FastAPI +from fastapi.middleware.cors import CORSMiddleware +from fastapi.responses import StreamingResponse from langchain_huggingface import HuggingFaceEmbeddings -from langfuse.langchain import CallbackHandler from langfuse import get_client +from langfuse.langchain import CallbackHandler from langgraph.checkpoint.redis.aio import AsyncRedisSaver from openai import APIConnectionError +from pydantic import BaseModel -from agents.graph import graph_builder, SynapticState, REDIS_URL -from memory.short_term import ShortTermMemory -from memory.long_term import LongTermMemory -from memory.manager import MemoryManager -from ingestion.stackoverflow_loader import EMBEDDING_MODEL, CONN_STR +import agents.nodes.memory_node as mem_module +import agents.nodes.orchestrator as orch_module +import agents.nodes.rag_agent as rag_agent_module +import agents.nodes.writer_node as writer_module +from agents.graph import REDIS_URL, SynapticState, graph_builder +from constants import SYSTEM_PROMPT from ingestion.stackoverflow_data_builder import ( - SODatasetBuilder, DATA_PATH, EVAL_IDS_PATH, + SODatasetBuilder, ) -from constants import SYSTEM_PROMPT -import agents.nodes.orchestrator as orch_module -import agents.nodes.memory_node as mem_module -import agents.nodes.writer_node as writer_module -import agents.nodes.rag_agent as rag_agent_module +from ingestion.stackoverflow_loader import CONN_STR, EMBEDDING_MODEL, IngestionPipeline +from llm import utility_llm +from memory.long_term import LongTermMemory +from memory.manager import MemoryManager +from memory.short_term import ShortTermMemory logger = logging.getLogger(__name__) app_graph = None -memory_manager: Optional[MemoryManager] = None +memory_manager: MemoryManager | None = None _metrics: dict[str, float] = { "total_queries": 0, "total_errors": 0, @@ -68,7 +66,7 @@ def chat_completion_chunk( created_at: int, model: str, delta: dict, - finish_reason: Optional[str] = None, + finish_reason: str | None = None, ) -> str: return f"data: {json.dumps({'id': completion_id, 'object': 'chat.completion.chunk', 'created': created_at, 'model': model, 'choices': [{'index': 0, 'delta': delta, 'finish_reason': finish_reason}]})}\n\n" @@ -133,7 +131,7 @@ class ChatCompletionRequest(BaseModel): model: str messages: list[ChatMessage] stream: bool = True - session_id: Optional[str] = None + session_id: str | None = None class IngestRequest(BaseModel): diff --git a/backend/memory/long_term.py b/backend/memory/long_term.py index ca3adbf..51be97a 100644 --- a/backend/memory/long_term.py +++ b/backend/memory/long_term.py @@ -1,6 +1,7 @@ -import asyncpg from typing import Any +import asyncpg + class LongTermMemory: TOP_K = 3 diff --git a/backend/memory/manager.py b/backend/memory/manager.py index 0121cf6..b01b412 100644 --- a/backend/memory/manager.py +++ b/backend/memory/manager.py @@ -1,8 +1,10 @@ -from .short_term import ShortTermMemory -from .long_term import LongTermMemory -from langchain_openai import ChatOpenAI import asyncio +from langchain_openai import ChatOpenAI + +from .long_term import LongTermMemory +from .short_term import ShortTermMemory + class MemoryManager: """ diff --git a/backend/memory/short_term.py b/backend/memory/short_term.py index 723941d..d6bdb14 100644 --- a/backend/memory/short_term.py +++ b/backend/memory/short_term.py @@ -1,16 +1,22 @@ import json -import tiktoken -import redis.asyncio as aioredis -from datetime import datetime, timezone +from datetime import UTC, datetime from typing import Any, cast +import redis.asyncio as aioredis +import tiktoken + +from context.engineer import AGENT_BUDGET + ENCODING = tiktoken.get_encoding("cl100k_base") class ShortTermMemory: - MAX_TURNS = 10 # pairs → up to 20 list entries + MAX_TURNS = 30 # pairs → up to 60 list entries; TOKEN_BUDGET is the real ceiling TTL_SECONDS = 86400 # 24-hour inactivity window (ARCHITECTURE.md §5.3) - TOKEN_BUDGET = 2000 # soft limit; oldest pairs dropped first if exceeded + # Soft limit; oldest pairs dropped first if exceeded. Must satisfy the + # smallest per-agent short_term_memory cap in AGENT_BUDGET, since this + # memory is loaded once before the graph routes to a specific agent. + TOKEN_BUDGET = min(b["short_term_memory"] for b in AGENT_BUDGET.values()) def __init__(self, redis_client: aioredis.Redis) -> None: self._r = redis_client @@ -31,7 +37,7 @@ async def append( entry = json.dumps({ "role": role, "content": content, - "timestamp": datetime.now(timezone.utc).isoformat(), + "timestamp": datetime.now(UTC).isoformat(), "agent": agent, }) async with self._r.pipeline() as pipe: diff --git a/backend/pyproject.toml b/backend/pyproject.toml new file mode 100644 index 0000000..1b9f6b8 --- /dev/null +++ b/backend/pyproject.toml @@ -0,0 +1,18 @@ +[tool.ruff] +line-length = 100 +target-version = "py312" + +[tool.ruff.lint] +select = [ + "E", # pycodestyle errors + "F", # pyflakes + "I", # isort (import sorting) + "UP", # pyupgrade + "B", # flake8-bugbear +] +ignore = [ + "E501", # line length handled by formatter, not a hard error +] + +[tool.ruff.lint.isort] +known-first-party = ["agents", "chain", "context", "ingestion", "memory", "retrieval", "tools"] diff --git a/backend/requirements.txt b/backend/requirements.txt index c5d1aba..07a3b49 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -58,3 +58,6 @@ PyGithub>=2.3.0 # backend/tools/github_tool.py # ── Utilities ───────────────────────────────────────────────────────────────── tqdm>=4.66.0 + +# ── Dev / Linting ───────────────────────────────────────────────────────────── +ruff>=0.8.0 diff --git a/backend/retrieval/chunks_retriever.py b/backend/retrieval/chunks_retriever.py index 0bdba34..ab6cca8 100644 --- a/backend/retrieval/chunks_retriever.py +++ b/backend/retrieval/chunks_retriever.py @@ -1,13 +1,14 @@ -from typing import Any, ClassVar, Optional +import os +from typing import Any, ClassVar + import numpy as np import psycopg -from pgvector.psycopg import register_vector, register_vector_async -from pydantic import ConfigDict -from langchain_huggingface import HuggingFaceEmbeddings +from dotenv import load_dotenv from langchain_core.documents import Document from langchain_core.retrievers import BaseRetriever -from dotenv import load_dotenv -import os +from langchain_huggingface import HuggingFaceEmbeddings +from pgvector.psycopg import register_vector, register_vector_async +from pydantic import ConfigDict from sentence_transformers import CrossEncoder load_dotenv() @@ -41,7 +42,7 @@ class ChunksRetriever(BaseRetriever): model_config = ConfigDict(arbitrary_types_allowed=True) embeddings: HuggingFaceEmbeddings - _encoder: ClassVar[Optional[CrossEncoder]] = None + _encoder: ClassVar[CrossEncoder | None] = None @classmethod def _get_encoder(cls) -> CrossEncoder: @@ -76,7 +77,9 @@ def _get_relevant_documents( re_ranked = self._get_encoder().predict([(query, row[0]) for row in rows]) - scored = sorted(zip(re_ranked, rows), key=lambda x: x[0], reverse=True) + scored = sorted( + zip(re_ranked, rows, strict=True), key=lambda x: x[0], reverse=True + ) top = [(score, r) for score, r in scored[:RETRIEVAL_TOP_K]] return [ @@ -108,7 +111,9 @@ async def _aget_relevant_documents( re_ranked = self._get_encoder().predict([(query, row[0]) for row in rows]) - scored = sorted(zip(re_ranked, rows), key=lambda x: x[0], reverse=True) + scored = sorted( + zip(re_ranked, rows, strict=True), key=lambda x: x[0], reverse=True + ) top = [(score, r) for score, r in scored[:RETRIEVAL_TOP_K]] return [