diff --git a/app/app.py b/app/app.py index 7b299a0..ce64a77 100644 --- a/app/app.py +++ b/app/app.py @@ -37,10 +37,31 @@ def notifications(ws): # ================ # DECORATORS # ================ + +import hashlib + +def make_cache_key(model, prompt): + raw = f"{model}:{prompt}" + return hashlib.sha256(raw.encode()).hexdigest() + +def get_cache(model, prompt): + key = make_cache_key(model, prompt) + db = SessionLocal() + entry = db.query(Cache).filter_by(key=key).first() + db.close() + return entry.response if entry else None + +def set_cache(model, prompt, response): + key = make_cache_key(model, prompt) + db = SessionLocal() + entry = Cache(key=key, model=model, response=response) + db.add(entry) + db.commit() + db.close() + + import structlog - logger = structlog.get_logger() - @app.before_request def log_request(): logger.info( @@ -1013,9 +1034,65 @@ def generate(): # return jsonify({"error": "Invalid JSON response from Ollama."}), 502 #output_text = ollama_data.get("response") or ollama_data.get("output") or "" #return jsonify({"model": model_name, "output": output_text}) - + + cached = get_cache(model, prompt) + if cached: + return jsonify({"model": model, "output": cached, "cached": True}) + result, used_model = call_ollama_with_fallback({"prompt": prompt}) - return jsonify({"model_used": used_model, "output": result.get("response")}) + output = result.get("response") + + set_cache(used_model, prompt, output) + + return jsonify({"model": used_model, "output": output, "cached": False}) + +@app.route("/api/multi", methods=["POST"]) +@login_required +@rate_limited +def multi(): + data = request.json + models = data.get("models", ["llama3.2", "gemma4"]) + prompt = data.get("prompt") + + results = {} + + for m in models: + try: + r = requests.post( + f"{app.config['OLLAMA_BASE_URL']}/api/generate", + json={"model": m, "prompt": prompt} + ) + r.raise_for_status() + results[m] = r.json().get("response") + except Exception as e: + results[m] = f"Error: {str(e)}" + + return jsonify(results) + +@app.route("/api/multi/stream") +@login_required +def multi_stream(): + models = request.args.get("models", "llama3.2,gemma4").split(",") + prompt = request.args.get("prompt") + + def generate(): + for m in models: + yield f"event: model\n" + yield f"data: {json.dumps({'model': m})}\n\n" + + url = f"{app.config['OLLAMA_BASE_URL']}/api/generate" + payload = {"model": m, "prompt": prompt, "stream": True} + + with requests.post(url, json=payload, stream=True) as r: + for line in r.iter_lines(): + if not line: + continue + data = json.loads(line.decode()) + token = data.get("response", "") + yield f"event: token\n" + yield f"data: {json.dumps({'model': m, 'token': token})}\n\n" + + return Response(generate(), mimetype="text/event-stream")