File size: 6,119 Bytes
a521353
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
"""FastAPI routes attached to gr.Server. The frontend talks ONLY to /api/*."""
import io, json, traceback, base64
from PIL import Image
from fastapi import UploadFile, File, Form, HTTPException
from fastapi.responses import JSONResponse, FileResponse
from fastapi.staticfiles import StaticFiles

import spaces

from .config import AUDIO_CACHE
from .model_loader import make_llm
from .grammar import load_grammar
from .prompt_builder import build_messages, new_session_meta
from .schema import ElysiumEnvelope, ElysiumResponse
from .hypergraph import persistence
from .hypergraph.engine import Hypergraph
from .tools.dispatcher import execute_all
from .tts.debate_sequencer import build_debate


# ─── Singletons ───
HG: Hypergraph = persistence.load()
GRAMMAR = load_grammar()


# ─── GPU-bound inference ───
@spaces.GPU(duration=120)
def _gpu_infer(messages: list, max_tokens: int = 4096) -> str:
    llm = make_llm()
    out = llm.create_chat_completion(
        messages=messages,
        max_tokens=max_tokens,
        temperature=0.7,
        grammar=GRAMMAR,          # strict JSON at sampling time
    )
    return out["choices"][0]["message"]["content"]


def _fallback_envelope(user_text: str, err: str) -> dict:
    meta = new_session_meta()
    resp = ElysiumResponse(
        session_id=meta["session_id"],
        timestamp_utc=meta["timestamp_utc"],
        interaction_type="SIMPLE_REPLY",
        direct_answer=f"(fallback) {err}",
    )
    return {"user_msg": user_text, "elysium_response": resp.model_dump()}


def attach(app):
    """Register all /api routes on the gr.Server FastAPI app."""

    # mount /audio for generated debate wavs
    app.mount("/audio", StaticFiles(directory=str(AUDIO_CACHE)), name="audio")

    @app.get("/api/health")
    async def health():
        return {"status": "ok",
                "nodes": HG.node_count(),
                "edges": HG.edge_count(),
                "grammar": GRAMMAR is not None}

    @app.get("/api/hypergraph")
    async def hypergraph():
        nodes, edges = [], []
        for i in HG.g.node_indexes():
            d = HG.g[i]
            nodes.append({"node_id": d["node_id"], "label": d["label"],
                          "node_type": d["node_type"], "payload": d.get("payload", {})})
        for s, t in HG.g.edge_list():
            d = HG.g.get_edge_data(s, t)
            edges.append({"edge_id": d["edge_id"],
                          "source_node_id": HG.g[s]["node_id"],
                          "target_node_id": HG.g[t]["node_id"],
                          "edge_type": d["edge_type"], "weight": d["weight"]})
        return {"nodes": nodes, "edges": edges,
                "node_count": HG.node_count(), "edge_count": HG.edge_count()}

    @app.post("/api/turn")
    async def turn(user_text: str = Form(""), image: UploadFile = File(None)):
        try:
            # 1. Load image if present
            img = None
            if image is not None:
                content = await image.read()
                if content:
                    try:
                        img = Image.open(io.BytesIO(content))
                    except Exception:
                        img = None

            # 2. Build messages with hypergraph context
            messages = build_messages(user_text, img, HG.context_summary())

            # 3. GPU inference (returns strict JSON)
            raw = _gpu_infer(messages)

            # 4. Parse
            try:
                envelope = ElysiumEnvelope.model_validate_json(raw)
            except Exception as parse_err:
                # try to extract any JSON object from raw
                try:
                    blob = json.loads(raw)
                    if "elysium_response" not in blob:
                        # wrap as direct_answer
                        meta = new_session_meta()
                        envelope = ElysiumEnvelope(
                            user_msg=user_text,
                            elysium_response=ElysiumResponse(
                                session_id=meta["session_id"],
                                timestamp_utc=meta["timestamp_utc"],
                                interaction_type="SIMPLE_REPLY",
                                direct_answer=str(blob)[:600]))
                    else:
                        envelope = ElysiumEnvelope.model_validate(blob)
                except Exception:
                    return JSONResponse(_fallback_envelope(user_text, f"parse_error: {parse_err}"))

            resp = envelope.elysium_response

            # 5. Apply hypergraph delta
            HG.apply_delta(resp.hypergraph_delta)
            persistence.save(HG)

            # 6. Execute tools
            tool_results = execute_all(resp.tool_calls) if resp.tool_calls else []

            # 7. Build audio drama if needed
            audio_url = None
            if resp.council_deliberation.debate_mode == "AUDIO_DRAMA" \
               and resp.council_deliberation.agent_outputs:
                try:
                    audio_url = build_debate(
                        [a.model_dump() for a in resp.council_deliberation.agent_outputs]
                    )
                except Exception as e:
                    print(f"[tts] debate failed: {e}")

            payload = envelope.model_dump()
            payload["_runtime"] = {
                "tool_results": tool_results,
                "audio_url": audio_url,
                "hypergraph": {"nodes": HG.node_count(), "edges": HG.edge_count()},
            }
            return JSONResponse(payload)

        except Exception as e:
            traceback.print_exc()
            return JSONResponse(_fallback_envelope(user_text, str(e)), status_code=200)

    @app.post("/api/reset")
    async def reset():
        """Wipe the hypergraph β€” start a new civilization."""
        global HG
        HG = Hypergraph()
        persistence.save(HG)
        return {"status": "reset"}