Files
astraea/app.py

193 lines
6.2 KiB
Python

# -*- coding: utf-8 -*-
"""Astraea — a constellation of Washington State family-law specialist agents.
Flask app + semantic RAG backend. LLM (granite4.2 RAG lane) and embeddings
(nomic-embed-text-v2-moe) run on the Ollama host (nightmare).
"""
import json
import re
import time
import urllib.request
from flask import Flask, jsonify, render_template, request
import rag
from agents import (
AGENTS, AGENT_BY_ID, OLLAMA_URL, RAG_MODEL, GENERAL_MODEL, DISCLAIMER,
)
app = Flask(__name__)
def _ollama_chat(model, messages, num_ctx=16384, num_predict=1024, temperature=0.1, timeout=180):
payload = {
"model": model,
"messages": messages,
"stream": False,
"think": False,
"options": {
"temperature": temperature,
"num_predict": num_predict,
"num_ctx": num_ctx,
},
}
req = urllib.request.Request(
f"{OLLAMA_URL}/api/chat",
data=json.dumps(payload).encode("utf-8"),
headers={"Content-Type": "application/json"},
)
opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
with opener.open(req, timeout=timeout) as r:
d = json.loads(r.read().decode("utf-8"))
msg = d.get("message", {})
content = msg.get("content") or msg.get("thinking") or ""
return content.strip()
def _strip_model_sources(text):
"""Remove a model-written trailing 'Sources:'/'References:' block (defensive)."""
text = re.split(r"\n\s*(?:Sources|References|Citations)\s*:\s*\n", text, flags=re.I)[0]
text = re.split(r"\n\s*Sources?\s*$", text, flags=re.I)[0]
return text.strip()
def _build_answer(agent, message, history):
top = rag.retrieve(message, agent["id"], top_k=5)
if not top:
return {
"answer": "I couldn't find relevant Washington law on that in my reference "
"library yet. Try rephrasing, or ask the Navigator to point you to "
"the right specialist.",
"citations": [],
"grounded": False,
}
context_blocks = []
citations = []
seen = set()
for i, (score, c) in enumerate(top, 1):
key = (c["source"], c["title"])
if key in seen:
continue
seen.add(key)
context_blocks.append(f"[{i}] ({c['source']} — {c['title']})\n{c['text']}")
citations.append({
"source": c["source"],
"title": c["title"],
"score": round(score, 3),
"snippet": c["text"][:260] + ("…" if len(c["text"]) > 260 else ""),
})
system = (
f"{agent['system']}\n\n"
"Rules:\n"
"- Answer the user's question DIRECTLY and concisely. Begin your answer immediately — "
"do NOT restate the question, do NOT narrate your reasoning, and do NOT say what you "
"are about to do.\n"
"- Answer using ONLY the reference documents provided below.\n"
"- Cite the RCW section (or source) inline for every legal claim, e.g. (RCW 26.09.030).\n"
"- If the documents do not contain the answer, say so clearly and suggest which "
"specialist or official resource to consult.\n"
"- Be practical and plain-English, specific to Washington State.\n"
"- Do NOT write a 'Sources' list at the end; cite inline only."
)
user = (
"REFERENCE DOCUMENTS:\n\n"
+ "\n\n".join(context_blocks)
+ f"\n\nUSER QUESTION: {message}"
)
messages = [{"role": "system", "content": system}]
if history:
for turn in history[-6:]:
if turn.get("role") in ("user", "assistant"):
messages.append({"role": turn["role"], "content": turn["content"]})
messages.append({"role": "user", "content": user})
try:
answer = _ollama_chat(RAG_MODEL, messages)
except Exception as e:
# Fall back to the general model if the RAG lane is unavailable.
try:
answer = _ollama_chat(GENERAL_MODEL, messages, num_ctx=16384)
except Exception:
return {
"answer": "The legal engine is temporarily unavailable. Please try again in a "
"moment. (Backend LLM could not be reached.)",
"citations": [],
"grounded": False,
}
answer = _strip_model_sources(answer)
return {
"answer": answer,
"citations": citations,
"grounded": True,
"model": RAG_MODEL,
}
@app.route("/")
def index():
agents_public = [
{k: a[k] for k in ("id", "name", "emoji", "tagline", "description", "accent")}
for a in AGENTS
]
return render_template("index.html", agents=agents_public, disclaimer=DISCLAIMER)
@app.route("/api/agents")
def api_agents():
return jsonify([
{k: a[k] for k in ("id", "name", "emoji", "tagline", "description", "accent")}
for a in AGENTS
])
@app.route("/api/chat", methods=["POST"])
def api_chat():
data = request.get_json(force=True, silent=True) or {}
agent_id = (data.get("agent_id") or "navigator").strip()
message = (data.get("message") or "").strip()
history = data.get("history") or []
agent = AGENT_BY_ID.get(agent_id)
if not agent:
return jsonify({"error": "Unknown agent"}), 400
if not message:
return jsonify({"error": "Message required"}), 400
if len(message) > 4000:
message = message[:4000]
t0 = time.time()
result = _build_answer(agent, message, history)
result["latency_ms"] = round((time.time() - t0) * 1000)
result["agent"] = agent_id
return jsonify(result)
@app.route("/api/health")
def health():
return jsonify({"ok": True, "agent_count": len(AGENTS), "index": rag.index_status()})
@app.route("/api/index/status")
def index_status():
return jsonify(rag.index_status())
@app.route("/api/reindex", methods=["POST"])
def reindex():
try:
idx = rag.build_index()
return jsonify({"ok": True, "chunks": len(idx)})
except Exception as e:
return jsonify({"ok": False, "error": str(e)}), 500
# Kick off the background index build once at import (before serving).
rag.start_background_index()
if __name__ == "__main__":
app.run(host="0.0.0.0", port=5000, threaded=True)