study-buddy / app /services /context_pack.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
11.4 kB
"""Context budgeting for ResearchMate agents.
This service is the boundary between retrieval/memory systems and prompts. It
turns raw Cognee recall plus raw paper chunks into a ranked, bounded context
pack so agents receive relevant evidence instead of an unbounded firehose.
"""
from __future__ import annotations
import re
from dataclasses import dataclass, field
from typing import Any, Iterable, Literal
from app.rag.context_composer import ContextComposer
from app.rag.models import ConsumerType, EvidenceRequest
ContextKind = Literal[
"student_profile",
"project_memory",
"source_material",
"conversation_history",
"recent_activity",
]
@dataclass
class ContextSection:
kind: ContextKind
title: str
content: str
priority: int
budget_chars: int
metadata: dict[str, Any] = field(default_factory=dict)
def render(self) -> str:
return f"{self.title}:\n{self.content.strip()}"
@dataclass
class ContextPack:
sections: list[ContextSection]
max_chars: int
def render(self) -> str:
parts: list[str] = []
remaining = self.max_chars
for section in self.sections:
rendered = section.render().strip()
if not rendered or remaining <= 0:
break
if len(rendered) > remaining:
rendered = _clip_at_boundary(rendered, remaining)
parts.append(rendered)
remaining -= len(rendered) + 2
return "\n\n".join(parts)
def render_memory_context(self) -> str:
return "\n\n".join(
section.render()
for section in self.sections
if section.kind in {"student_profile", "project_memory", "recent_activity"}
).strip()
def render_source_context(self) -> str:
return "\n\n".join(
section.content
for section in self.sections
if section.kind == "source_material"
).strip()
class ContextPackService:
def __init__(self, composer: ContextComposer | None = None) -> None:
self._composer = composer or ContextComposer()
async def build(
self,
*,
agent_id: str,
query: str,
project_id: str,
paper_chunks: list[dict[str, Any]] | None = None,
history: list[dict[str, Any]] | None = None,
recent_summaries: list[str] | None = None,
max_chars: int = 9000,
) -> ContextPack:
intent = _classify_intent(query)
budgets = _budgets_for(intent, agent_id, max_chars)
try:
consumer = ConsumerType(agent_id)
except ValueError:
consumer = ConsumerType.CHAT
project_items, student_items = await self._composer.recall_memory(
EvidenceRequest(
query=query or "relevant project and student context",
consumer=consumer,
project_id=project_id,
token_budget=max(256, max_chars // 4),
project_memory_policy="relevant",
student_memory_policy=(
"live" if agent_id in {"brain", "project_graph"} else "cached"
),
)
)
profile_raw = "\n".join(item.statement for item in student_items)
project_raw = "\n".join(item.statement for item in project_items)
sections = [
ContextSection(
kind="source_material",
title="SOURCE MATERIAL",
content=_build_source_context(paper_chunks or [], query, budgets["source_material"]),
priority=_priority(intent, "source_material"),
budget_chars=budgets["source_material"],
),
ContextSection(
kind="student_profile",
title="STUDENT PROFILE",
content=_compress_text(profile_raw, query, budgets["student_profile"], preserve_profile_facts=True),
priority=_priority(intent, "student_profile"),
budget_chars=budgets["student_profile"],
),
ContextSection(
kind="project_memory",
title="PROJECT MEMORY",
content=_compress_text(project_raw, query, budgets["project_memory"]),
priority=_priority(intent, "project_memory"),
budget_chars=budgets["project_memory"],
),
ContextSection(
kind="conversation_history",
title="RECENT CONVERSATION",
content=_build_history_context(history or [], query, budgets["conversation_history"]),
priority=_priority(intent, "conversation_history"),
budget_chars=budgets["conversation_history"],
),
ContextSection(
kind="recent_activity",
title="RECENT ACTIVITY",
content=_compress_text("\n".join(recent_summaries or []), query, budgets["recent_activity"]),
priority=_priority(intent, "recent_activity"),
budget_chars=budgets["recent_activity"],
),
]
sections = [s for s in sections if s.content.strip()]
sections.sort(key=lambda s: s.priority)
return ContextPack(sections=sections, max_chars=max_chars)
def _classify_intent(query: str) -> str:
lowered = (query or "").lower()
if any(term in lowered for term in ("my name", "who am i", "call me", "preference", "prefer", "learning style")):
return "identity"
if any(term in lowered for term in ("paper", "source", "chunk", "figure", "equation", "theorem", "method", "explain")):
return "paper"
return "study"
def _budgets_for(intent: str, agent_id: str, max_chars: int) -> dict[ContextKind, int]:
if intent == "identity":
shares = {
"student_profile": 0.50,
"project_memory": 0.20,
"recent_activity": 0.15,
"conversation_history": 0.10,
"source_material": 0.05,
}
elif intent == "paper":
shares = {
"source_material": 0.58,
"student_profile": 0.12,
"project_memory": 0.16,
"conversation_history": 0.10,
"recent_activity": 0.04,
}
else:
shares = {
"source_material": 0.38,
"project_memory": 0.24,
"student_profile": 0.18,
"conversation_history": 0.12,
"recent_activity": 0.08,
}
if agent_id == "pair_buddy":
shares["recent_activity"] = max(shares["recent_activity"], 0.14)
shares["conversation_history"] = max(shares["conversation_history"], 0.14)
return {kind: max(120, int(max_chars * share)) for kind, share in shares.items()} # type: ignore[return-value]
def _priority(intent: str, kind: ContextKind) -> int:
orders = {
"identity": ["student_profile", "project_memory", "recent_activity", "conversation_history", "source_material"],
"paper": ["source_material", "student_profile", "project_memory", "conversation_history", "recent_activity"],
"study": ["source_material", "project_memory", "student_profile", "conversation_history", "recent_activity"],
}
return orders.get(intent, orders["study"]).index(kind)
def _build_source_context(chunks: list[dict[str, Any]], query: str, limit: int) -> str:
ranked = sorted(chunks, key=lambda chunk: _score_text(str(chunk.get("text", "")), query), reverse=True)
parts = []
remaining = limit
for i, chunk in enumerate(ranked):
source = chunk.get("source") or chunk.get("filename") or "source"
if _terms(query) and _score_text(str(chunk.get("text", "")), query) <= 0:
continue
text = _compress_text(str(chunk.get("text", "")), query, min(remaining, max(220, limit // 2)))
if not text:
continue
part = f"[{source}]\n{text}"
if len(part) > remaining:
part = _clip_at_boundary(part, remaining)
parts.append(part)
remaining -= len(part) + 2
if remaining <= 120 or i >= 4:
break
return "\n\n".join(parts)
def _build_history_context(history: list[dict[str, Any]], query: str, limit: int) -> str:
recent = history[-6:]
text = "\n".join(f"{item.get('role', 'turn')}: {item.get('content', '')}" for item in recent)
return _compress_text(text, query, limit)
def _compress_text(text: str, query: str, limit: int, preserve_profile_facts: bool = False) -> str:
cleaned = _normalise(text)
if not cleaned or limit <= 0:
return ""
sentences = _split_units(cleaned)
selected: list[str] = []
if preserve_profile_facts:
for unit in sentences:
lowered = unit.lower()
if any(term in lowered for term in ("name", "call me", "preferred", "prefer", "learning style")):
selected.append(unit)
scores = {unit: _score_text(unit, query) for unit in sentences}
has_relevant_units = bool(_terms(query)) and any(score > 0 for score in scores.values())
ranked = sorted(sentences, key=lambda unit: scores[unit], reverse=True)
for unit in ranked:
if has_relevant_units and scores[unit] <= 0:
continue
if unit not in selected:
selected.append(unit)
return _join_with_budget(_dedupe(selected), limit)
def _score_text(text: str, query: str) -> int:
terms = _terms(query)
lowered = text.lower()
score = sum(1 for term in terms if term in lowered)
if any(term in lowered for term in ("name", "preferred", "prefer", "confused", "struggle")):
score += 2
return score
def _terms(query: str) -> set[str]:
return {term for term in re.findall(r"[a-zA-Z][a-zA-Z0-9_-]{2,}", (query or "").lower()) if term not in _STOPWORDS}
def _split_units(text: str) -> list[str]:
parts = re.split(r"(?<=[.!?])\s+|\n+", text)
return [part.strip() for part in parts if part.strip()]
def _dedupe(items: Iterable[str]) -> list[str]:
seen = set()
result = []
for item in items:
key = re.sub(r"\s+", " ", item.lower()).strip()
if key in seen:
continue
seen.add(key)
result.append(item)
return result
def _join_with_budget(items: list[str], limit: int) -> str:
selected = []
remaining = limit
for item in items:
if remaining <= 0:
break
if len(item) > remaining:
selected.append(_clip_at_boundary(item, remaining))
break
selected.append(item)
remaining -= len(item) + 1
return "\n".join(selected).strip()
def _clip_at_boundary(text: str, limit: int) -> str:
if len(text) <= limit:
return text
if limit <= 20:
return text[:limit]
clipped = text[: limit - 15].rstrip()
boundary = max(clipped.rfind("."), clipped.rfind("\n"), clipped.rfind(" "))
if boundary > limit // 2:
clipped = clipped[:boundary].rstrip()
return clipped + "\n...[truncated]"
def _normalise(text: str) -> str:
return re.sub(r"\s+", " ", text or "").strip()
_STOPWORDS = {
"about",
"after",
"again",
"from",
"have",
"into",
"paper",
"source",
"that",
"their",
"there",
"this",
"what",
"when",
"where",
"which",
"with",
}