Commit e867818e authored by Kostas Chartsias's avatar Kostas Chartsias
Browse files

endpoint for SSE support

parent 6fae1d4b
Loading
Loading
Loading
Loading
+30 −1
Original line number Diff line number Diff line
import logging
from quart import request, jsonify
from quart import request, jsonify, Response
from AI_agent.sub_agents.groq_agent import create_groq_agent
from AI_agent.utils import stream_agent_response

def register_groq_routes(app):

@@ -21,3 +22,31 @@ def register_groq_routes(app):
            return jsonify({"error": str(e)}), 500
        finally:
            await agent.close()

    @app.route("/groq-mcp/stream", methods=["GET"])
    async def groq_stream():
        query = request.args.get("query")

        if not query:
            return jsonify({"error": "No query provided"}), 400

        agent = await create_groq_agent()

        try:
            generator = await stream_agent_response(agent, query)
            return Response(
                generator(),
                content_type="text/event-stream; charset=utf-8",
                headers={
                    "Cache-Control": "no-cache",
                    "Connection": "keep-alive",
                    "X-Accel-Buffering": "no",
                    "Access-Control-Allow-Origin": "*",
                },
            )
        except Exception as e:
            logging.error("Error in groq_stream", exc_info=True)
            await agent.close()
            return jsonify({"error": str(e)}), 500

AI_agent/utils.py

0 → 100644
+43 −0
Original line number Diff line number Diff line
import asyncio
import json
import logging

# --- SSE Helper ---
def sse_event(data, event="message", id=None, retry=None):
    lines = []
    if id is not None:
        lines.append(f"id: {id}")
    if event is not None:
        lines.append(f"event: {event}")
    if retry is not None:
        lines.append(f"retry: {retry}")

    if isinstance(data, str):
        for line in data.splitlines():
            lines.append(f"data: {line}")
    else:
        json_data = json.dumps(data)
        for line in json_data.splitlines():
            lines.append(f"data: {line}")

    return "\n".join(lines) + "\n\n"


# --- Common SSE Stream Helper ---
async def stream_agent_response(agent, query):
    chunk_size = 50
    async def generator():
        event_id = 0
        try:
            result = await agent.run(query)
            for i in range(0, len(result), chunk_size):
                yield sse_event({"chunk": result[i:i + chunk_size]}, id=event_id)
                event_id += 1
                await asyncio.sleep(0.01)
            yield sse_event({"done": True}, event="complete", id=event_id)
        except Exception as e:
            yield sse_event({"error": str(e)}, event="error", id=event_id)
        finally:
            await agent.close()
            logging.info("Agent closed after streaming")
    return generator