lightning / gen.py
sharktide's picture
Update gen.py
7bfd539 verified
Raw
History Blame
36.2 kB
import os
import base64
import random
import httpx
from urllib.parse import quote
from fastapi import APIRouter, Request, HTTPException, Header
from fastapi.responses import Response, JSONResponse, StreamingResponse
import re
from typing import Optional, Any
import json
from helper.assets import (
save_base64_image,
cleanup_image,
is_base64_image,
)
import asyncio
from helper.ratelimit import (
enforce_rate_limit,
resolve_rate_limit_identity,
check_audio_rate_limit,
check_video_rate_limit,
check_image_rate_limit,
MAX_CHAT_PROMPT_BYTES,
MAX_CHAT_PROMPT_CHARS,
MAX_GROQ_PROMPT_BYTES,
MAX_GROQ_PROMPT_CHARS,
MAX_MEDIA_PROMPT_BYTES,
MAX_MEDIA_PROMPT_CHARS,
extract_user_text,
calculate_messages_size,
normalize_prompt_value,
enforce_prompt_size,
resolve_bound_subject,
get_usage_snapshot_for_subject,
)
from helper.keywords import *
from uuid import uuid4
from time import time
from typing import Dict, List, Optional, Tuple
router = APIRouter(prefix="/gen")
PKEY = os.getenv("POLLINATIONS_KEY", "")
PKEY2 = os.getenv("POLLINATIONS2_KEY", "")
PKEY3 = os.getenv("POLLINATIONS3_KEY", "")
AIRFORCE_KEY = os.getenv("AIRFORCE")
AIRFORCE_VIDEO_MODEL = "grok-imagine-video"
AIRFORCE_API_URL = "https://api.airforce/v1/images/generations"
valid_ratios = {"3:2", "2:3", "1:1", "", None}
ratios = {"3:2", "2:3", "1:1"}
valid_modes = {"normal", "fun", "", None}
modes = {"normal", "fun"}
MODEL_MAP = {
"llama-3.1-8b-instant": "Meta Llama 3.1 8B Instant",
"gpt-4o-mini": "OpenAI GPT 4o Mini",
"nemotron-3-super": "NVIDIA Nemotron 3 Super",
"openai/gpt-oss-120b": "OpenAI GPT-OSS 120B",
"openai/gpt-oss-20b": "OpenAI GPT-OSS 20B",
"qwen-3-235b-a22b-instruct-2507": "Qwen3 Instruct",
"llama-3.3-70b-versatile": "Meta Llama 3.3 70B Versatile",
"meta-llama/llama-4-scout-17b-16e-instruct": "Meta Llama 4 Scout",
}
FALLBACK_MODEL = "meta-llama/llama-4-scout-17b-16e-instruct"
FALLBACK_PROVIDER = "groq"
# ──────────────────────────────────────────────
# CENTRAL ROUTING LOGIC
# ──────────────────────────────────────────────
def route_chat(
messages: List[Dict[str, Any]],
uses_tools: bool = False,
) -> Tuple[str, str]:
"""
Inspect messages and return (chosen_model, provider).
This is the single source of truth for model selection.
No API calls, no side-effects — pure routing logic.
"""
total_chars, total_bytes = calculate_messages_size(messages)
prompt_text = extract_user_text(messages)
long_context = is_long_context(messages)
code_present = contains_code(prompt_text)
math_heavy = is_math_heavy(prompt_text)
structured_task = is_structured_task(prompt_text)
multi_q = multiple_questions(prompt_text)
code_heavy = is_code_heavy(prompt_text, code_present, long_context)
has_images = contains_images(messages)
score = 0
if long_context: score += 3
if math_heavy: score += 3
if structured_task: score += 2
if code_present: score += 2
if multi_q: score += 1
for kw in REASONING_KEYWORDS:
if kw in prompt_text:
score += 1
score = min(score, 10)
# ── multimodal fast-path ──────────────────
if has_images:
return "gpt-4o-mini", "navy vision"
# ── tool-use branch ──────────────────────
if uses_tools:
if score >= 6:
return "nemotron-3-super", "navy"
if score >= 4:
return "openai/gpt-oss-120b", "groq"
return "openai/gpt-oss-20b", "groq"
# ── code branch ──────────────────────────
if code_present:
if code_heavy and score >= 6:
return "o3-mini", "navy"
if score >= 4:
return "llama-3.3-70b-versatile", "groq"
# ── general reasoning branch ─────────────
if score >= 6:
return "sonar", "navy"
if score >= 4:
return "meta-llama/llama-4-scout-17b-16e-instruct", "groq"
# ── default ──────────────────────────────
chosen_model, provider = "llama-3.1-8b-instant", "groq"
# Groq context-size guard — promote to navy if too large
if provider == "groq" and (
total_chars > MAX_GROQ_PROMPT_CHARS or total_bytes > MAX_GROQ_PROMPT_BYTES
):
return "gpt-4o-mini", "navy"
return chosen_model, provider
def _log_routing(
chosen_model: str,
provider: str,
messages: List[Dict[str, Any]],
uses_tools: bool,
) -> None:
prompt_text = extract_user_text(messages)
long_context = is_long_context(messages)
code_present = contains_code(prompt_text)
math_heavy = is_math_heavy(prompt_text)
structured_task = is_structured_task(prompt_text)
multi_q = multiple_questions(prompt_text)
has_images = contains_images(messages)
print(
f"\n[ADVANCED ROUTER]\n"
f" Uses tools: {uses_tools}\n"
f" Long context: {long_context}\n"
f" Code present: {code_present}\n"
f" Math heavy: {math_heavy}\n"
f" Structured: {structured_task}\n"
f" Multi-question:{multi_q}\n"
f" Has images: {has_images}\n"
f" → Selected: {chosen_model} ({provider})\n"
)
# ──────────────────────────────────────────────
# CENTRAL HTTP CALL
# ──────────────────────────────────────────────
def _get_provider_url_and_key(provider: str) -> Tuple[str, str]:
"""Return (url, api_key) for the given provider, raising on misconfiguration."""
if provider == "groq":
keys = [k.strip() for k in os.getenv("GROQ_KEY", "").split(",") if k.strip()]
if not keys:
raise HTTPException(500, "Missing GROQ_KEY(s)")
return "https://api.groq.com/openai/v1/chat/completions", random.choice(keys)
if provider == "cerebras":
keys = [k.strip() for k in os.getenv("CER_KEY", "").split(",") if k.strip()]
if not keys:
raise HTTPException(500, "Missing CER_KEY(s)")
return "https://api.cerebras.ai/v1/chat/completions", random.choice(keys)
if provider == "navy vision":
keys = [k.strip() for k in os.getenv("NAVY_KEY", "").split(",") if k.strip()]
if not keys:
raise HTTPException(500, "Missing NAVY_KEY(s)")
return "https://api.navy/v1/chat/completions", random.choice(keys)
if provider == "navy":
keys = [k.strip() for k in os.getenv("NAVY_TEXT_ONLY", "").split(",") if k.strip()]
if not keys:
raise HTTPException(500, "Missing NAVY_TEXT_ONLY key(s)")
return "https://api.navy/v1/chat/completions", random.choice(keys)
raise HTTPException(500, f"Unknown provider: {provider!r}")
async def call_chat_completions(
messages: List[Dict[str, Any]],
model: str,
provider: str,
extra_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""
Non-streaming chat-completions call.
Returns the full upstream JSON payload.
Raises HTTPException on upstream errors.
"""
url, api_key = _get_provider_url_and_key(provider)
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
body = {"model": model, "messages": messages, "stream": False}
if extra_body:
body.update(extra_body)
async with httpx.AsyncClient(timeout=None) as client:
r = await client.post(url, json=body, headers=headers)
if r.status_code != 200:
raise HTTPException(status_code=r.status_code, detail=r.text[:1000])
return r.json()
def _extract_text_from_response(data: Dict[str, Any]) -> str:
try:
return data["choices"][0]["message"]["content"] or ""
except Exception:
return ""
def _extract_usage(data: Dict[str, Any]) -> Tuple[int, int]:
usage = data.get("usage", {})
return usage.get("prompt_tokens", 0), usage.get("completion_tokens", 0)
# ──────────────────────────────────────────────
# HELPER: image generation
# ──────────────────────────────────────────────
def is_cinematic_image_prompt(prompt: str) -> bool:
for kw in CREATIVE_KEYWORDS:
if kw in prompt.lower():
return True
return False
# ──────────────────────────────────────────────
# IMAGE GENERATION
# ──────────────────────────────────────────────
@router.post("/image")
@router.get("/image/{prompt}")
async def generate_image(
request: Request,
prompt: str = None,
authorization: str = Header(None),
x_client_id: str = Header(None),
):
timeout = httpx.Timeout(300.0, read=300.0)
if prompt is None:
payload = await request.json()
prompt = payload.get("prompt")
mode = payload.get("mode")
image_urls = payload.get("image_urls")
else:
mode = request.query_params.get("mode")
image_urls = request.query_params.getlist("image_urls")
prompt = normalize_prompt_value(prompt, "prompt")
enforce_prompt_size(prompt, MAX_MEDIA_PROMPT_CHARS, MAX_MEDIA_PROMPT_BYTES, "Image prompt")
await check_image_rate_limit(request, authorization, x_client_id)
chosen_model = "zimage"
if is_cinematic_image_prompt(prompt):
chosen_model = "flux"
if isinstance(mode, str):
m = mode.strip().lower()
if m == "fantasy":
chosen_model = "flux"
elif m == "realistic":
chosen_model = "zimage"
has_input_image = bool(image_urls)
temp_assets = []
if has_input_image:
chosen_model = "klein"
params = {"model": chosen_model, "key": PKEY2}
if has_input_image:
processed = []
for img in image_urls[:2]:
if is_base64_image(img):
image_id = save_base64_image(img)
temp_assets.append(image_id)
served = f"{request.base_url}asset-cdn/assets/{image_id}"
processed.append(served)
else:
processed.append(img)
params["image"] = "|".join(processed)
encoded_prompt = quote(prompt, safe="")
query = "&".join(f"{k}={quote(str(v), safe='')}" for k, v in params.items())
url = f"https://gen.pollinations.ai/image/{encoded_prompt}?{query}"
try:
async with httpx.AsyncClient(timeout=timeout) as client:
resp = await client.get(url)
finally:
for aid in temp_assets:
cleanup_image(aid)
if resp.status_code != 200:
raise HTTPException(500, f"Pollinations error: {resp.status_code}")
return Response(content=resp.content, media_type="image/jpeg")
# ──────────────────────────────────────────────
# SFX GENERATION
# ──────────────────────────────────────────────
@router.get("/sfx/{prompt}")
@router.post("/sfx")
async def gensfx(
request: Request,
prompt: str = None,
authorization: str = Header(None),
x_client_id: str = Header(None),
):
if prompt is None:
payload = await request.json()
prompt = payload.get("prompt")
prompt = normalize_prompt_value(prompt, "prompt")
enforce_prompt_size(prompt, MAX_MEDIA_PROMPT_CHARS, MAX_MEDIA_PROMPT_BYTES, "Audio prompt")
await check_audio_rate_limit(request, authorization, x_client_id)
url = f"https://gen.pollinations.ai/audio/{prompt}?model=acestep&key={PKEY}"
async with httpx.AsyncClient(timeout=None) as client:
resp = await client.get(url)
if resp.status_code != 200:
return JSONResponse(
status_code=resp.status_code,
content={"success": False, "error": "Upstream music/sfx generation failed"},
)
return Response(resp.content, media_type="audio/mpeg")
# ──────────────────────────────────────────────
# TTS GENERATION
# ──────────────────────────────────────────────
@router.get("/tts/{prompt}")
@router.post("/tts")
async def gentts(
request: Request,
prompt: str = None,
authorization: str = Header(None),
x_client_id: str = Header(None),
):
if prompt is None:
payload = await request.json()
prompt = payload.get("prompt")
prompt = normalize_prompt_value(prompt, "prompt")
enforce_prompt_size(prompt, MAX_MEDIA_PROMPT_CHARS, MAX_MEDIA_PROMPT_BYTES, "Audio prompt")
await check_audio_rate_limit(request, authorization, x_client_id)
url = f"https://gen.pollinations.ai/audio/{prompt}?key={PKEY3}"
async with httpx.AsyncClient(timeout=None) as client:
resp = await client.get(url)
if resp.status_code != 200:
return JSONResponse(
status_code=resp.status_code,
content={"success": False, "error": "Upstream audio generation failed"},
)
return Response(resp.content, media_type="audio/mpeg")
# ──────────────────────────────────────────────
# VIDEO GENERATION (Pollinations)
# ──────────────────────────────────────────────
@router.get("/video/{prompt}")
@router.post("/video")
@router.head("/video")
async def genvideo(
request: Request,
prompt: str = None,
authorization: str = Header(None),
x_client_id: str = Header(None),
):
if request.method == "HEAD":
return Response(
status_code=200,
headers={
"Y-prompt": "string — required. The text prompt used to generate the video.",
"Y-ratio": "string — optional. Aspect ratio of the output video.",
"Y-ratio-values": "3:2,2:3,1:1",
"Y-ratio-default": "3:2",
"Y-mode": "string — optional. Controls generation style.",
"Y-mode-values": "normal,fun",
"Y-mode-default": "normal",
"Y-duration": "integer — optional. Duration in seconds (1–10).",
"Y-duration-default": "5",
"Y-image_urls": "array<string> — optional. Up to 2 image URLs for conditioning.",
"Y-image_urls-max": "2",
"Y-response_format": "video/mp4",
"Y-model": "grok-video",
},
)
aspectRatio = "3:2"
inputMode = "normal"
duration = 5
image_urls = None
if prompt is None:
user_body = await request.json()
prompt = user_body.get("prompt")
ratio = user_body.get("ratio")
mode = user_body.get("mode")
image_urls = user_body.get("image_urls")
duration = user_body.get("duration", 5)
if ratio not in valid_ratios:
raise HTTPException(400, f"Invalid aspect ratio '{ratio}'. Must be one of 3:2, 2:3, or 1:1.")
if ratio in ratios:
aspectRatio = ratio
if mode not in valid_modes:
raise HTTPException(400, f"Invalid mode '{mode}'. Must be 'normal' or 'fun'.")
if mode in modes:
inputMode = mode
if image_urls:
if not isinstance(image_urls, list):
raise HTTPException(400, "image_urls must be a list")
if len(image_urls) > 2:
raise HTTPException(400, "You may provide at most two image URLs")
try:
duration = max(1, min(10, int(duration)))
except (TypeError, ValueError):
duration = 5
prompt = normalize_prompt_value(prompt, "prompt")
enforce_prompt_size(prompt, MAX_MEDIA_PROMPT_CHARS, MAX_MEDIA_PROMPT_BYTES, "Video prompt")
await check_video_rate_limit(request, authorization, x_client_id)
RATIO_MAP = {"3:2": "16:9", "2:3": "9:16", "1:1": "9:16"}
pollinations_ratio = RATIO_MAP.get(aspectRatio, "16:9")
encoded_prompt = quote(prompt, safe="")
params = {
"model": "ltx-2",
"duration": duration,
"aspectRatio": pollinations_ratio,
"seed": -1,
}
temp_assets = []
if image_urls:
processed_urls = []
for img in image_urls[:2]:
if is_base64_image(img):
image_id = save_base64_image(img)
temp_assets.append(image_id)
served_url = f"{request.base_url}asset-cdn/assets/{image_id}"
processed_urls.append(served_url)
else:
processed_urls.append(img)
params["image"] = "|".join(processed_urls)
if inputMode == "fun":
params["enhance"] = "true"
query_string = "&".join(f"{k}={quote(str(v), safe='')}" for k, v in params.items())
url = f"https://gen.pollinations.ai/image/{encoded_prompt}?{query_string}&key={PKEY}"
print(f"[VIDEO GEN] Pollinations URL: {url}")
resp = None
try:
async with httpx.AsyncClient(timeout=600) as client:
resp = await client.get(url)
finally:
for aid in temp_assets:
cleanup_image(aid)
if resp is None:
raise HTTPException(502, "Video generation request failed")
if resp.status_code != 200:
body_text = ""
try:
body_text = resp.text
except Exception:
pass
return JSONResponse(
status_code=resp.status_code,
content={
"success": False,
"error": "Upstream video generation failed",
"status_code": resp.status_code,
"message": body_text[:1000],
},
)
if not resp.content:
raise HTTPException(502, "Pollinations returned empty response")
return Response(
content=resp.content,
media_type="video/mp4",
headers={
"Content-Length": str(len(resp.content)),
"Accept-Ranges": "bytes",
},
)
# ──────────────────────────────────────────────
# VIDEO GENERATION (Airforce)
# ──────────────────────────────────────────────
@router.get("/video/airforce/{prompt}")
@router.post("/video/airforce")
async def genvideo_airforce(
request: Request,
prompt: str = None,
authorization: str = Header(None),
x_client_id: str = Header(None),
):
if request.method == "HEAD":
return Response(
status_code=200,
headers={
"Y-prompt": "string — required. The text prompt used to generate the video.",
"Y-ratio": "string — optional. Aspect ratio of the output video.",
"Y-ratio-values": "3:2,2:3,1:1",
"Y-ratio-default": "3:2",
"Y-mode": "string — optional. Controls generation style.",
"Y-mode-values": "normal,fun",
"Y-mode-default": "normal",
"Y-duration": "integer — optional. Duration in seconds.",
"Y-duration-default": "5",
"Y-image_urls": "array<string> — optional. Up to 2 image URLs for conditioning.",
"Y-image_urls-max": "2",
"Y-response_format": "video/mp4",
"Y-model": "grok-imagine-video",
},
)
aspectRatio = "3:2"
inputMode = "normal"
image_urls = None
if prompt is None:
user_body = await request.json()
prompt = user_body.get("prompt")
ratio = user_body.get("ratio")
mode = user_body.get("mode")
image_urls = user_body.get("image_urls")
if ratio not in valid_ratios:
raise HTTPException(400, f"Invalid aspect ratio {ratio}. Must be one of 3:2, 2:3, or 1:1. Default is 3:2")
if ratio in ratios:
aspectRatio = ratio
if mode not in valid_modes:
raise HTTPException(400, f"Invalid mode {mode}. Must be 'normal' or 'fun'. Default is normal")
if mode in modes:
inputMode = mode
if image_urls:
if not isinstance(image_urls, list):
raise HTTPException(400, "image_urls must be a list")
if len(image_urls) > 2:
raise HTTPException(400, "You may provide at most two image URLs")
prompt = normalize_prompt_value(prompt, "prompt")
enforce_prompt_size(prompt, MAX_MEDIA_PROMPT_CHARS, MAX_MEDIA_PROMPT_BYTES, "Video prompt")
await check_video_rate_limit(request, authorization, x_client_id)
payload = {
"model": AIRFORCE_VIDEO_MODEL,
"prompt": prompt,
"n": 1,
"size": "1024x1024",
"response_format": "b64_json",
"sse": False,
"mode": inputMode,
"aspectRatio": aspectRatio,
}
if image_urls:
payload["image_urls"] = image_urls
async with httpx.AsyncClient(timeout=600) as client:
resp = await client.post(
AIRFORCE_API_URL,
headers={"Authorization": f"Bearer {AIRFORCE_KEY}", "Content-Type": "application/json"},
json=payload,
)
if resp.status_code != 200:
return JSONResponse(status_code=resp.status_code, content=resp.json())
if not resp.content:
raise HTTPException(502, "api.airforce returned empty response")
try:
result = resp.json()
b64_video = result["data"][0]["b64_json"]
except Exception:
raise HTTPException(502, f"Invalid api.airforce response: {resp.text[:500]}")
if not b64_video:
raise HTTPException(502, "Airforce returned empty b64_json")
video_bytes = base64.b64decode(b64_video)
return Response(
content=video_bytes,
media_type="video/mp4",
headers={
"Content-Length": str(len(video_bytes)),
"Accept-Ranges": "bytes",
},
)
# ──────────────────────────────────────────────
# CHAT COMPLETIONS (/gen/chat/completions)
# ──────────────────────────────────────────────
async def _check_chat_rate_limit(
request: Request,
authorization: Optional[str],
client_id: Optional[str] = None,
):
return await enforce_rate_limit(request, authorization, "cloudChatDaily", client_id)
@router.post("/chat/completions")
async def generate_text(
request: Request,
authorization: Optional[str] = Header(None),
x_client_id: Optional[str] = Header(None),
):
body = await request.json()
messages = body.get("messages", [])
if not isinstance(messages, list) or len(messages) == 0:
raise HTTPException(400, "messages[] is required")
uses_tools = (
"tools" in body and isinstance(body["tools"], list) and len(body["tools"]) > 0
) or ("tool_choice" in body and body["tool_choice"] not in [None, "none"])
chosen_model, provider = route_chat(messages, uses_tools=uses_tools)
_log_routing(chosen_model, provider, messages, uses_tools)
await _check_chat_rate_limit(request, authorization, x_client_id)
body["model"] = chosen_model
stream = body.get("stream", False)
url, api_key = _get_provider_url_and_key(provider)
headers = {"Authorization": f"Bearer {api_key}"}
if stream:
body["stream"] = True
async def stream_fallback(client: httpx.AsyncClient):
fallback_body = {
"model": FALLBACK_MODEL,
"messages": body["messages"],
"stream": True,
}
fb_url, fb_key = _get_provider_url_and_key(FALLBACK_PROVIDER)
fb_headers = {"Authorization": f"Bearer {fb_key}"}
print("[FALLBACK] Starting Groq fallback stream")
async with client.stream("POST", fb_url, json=fallback_body, headers=fb_headers) as r:
if r.status_code >= 400:
err = (await r.aread()).decode("utf-8", errors="replace")
yield f'data: {{"error": "Fallback provider failed: {err[:500]}"}}\n\n'
return
async for line in r.aiter_lines():
if not line:
yield "\n"
continue
yield (line if line.startswith("data:") else f"data: {line}\n\n") + "\n"
async def stream_primary(client: httpx.AsyncClient):
try:
async with client.stream("POST", url, json=body, headers=headers) as r:
if r.status_code >= 400:
print("[STREAM FALLBACK] Primary provider failed → switching to fallback")
async for chunk in stream_fallback(client):
yield chunk
return
async for line in r.aiter_lines():
if not line:
yield "\n"
continue
if line.startswith("data:"):
try:
obj = json.loads(line[5:].strip())
if isinstance(obj, dict) and isinstance(obj.get("error"), dict):
async for chunk in stream_fallback(client):
yield chunk
return
except Exception:
pass
yield line + "\n"
except Exception as e:
print(f"[STREAM ERROR] {e}")
async for chunk in stream_fallback(client):
yield chunk
async def event_generator():
sent_metadata = False
async with httpx.AsyncClient(timeout=None) as client:
async for chunk in stream_primary(client):
if not sent_metadata:
meta = {"router_metadata": {"model_name": MODEL_MAP.get(chosen_model, chosen_model)}}
yield f"data: {json.dumps(meta)}\n\n"
sent_metadata = True
yield chunk
return StreamingResponse(
event_generator(),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "Connection": "keep-alive", "X-Accel-Buffering": "no"},
)
# ── non-streaming ─────────────────────────
async with httpx.AsyncClient(timeout=None) as client:
r = await client.post(url, json=body, headers=headers)
# navy-vision fallback
if provider == "navy vision" and r.status_code >= 400:
print("[FALLBACK] Navy vision failed — switching to fallback")
fb_url, fb_key = _get_provider_url_and_key(FALLBACK_PROVIDER)
fallback_body = dict(body)
fallback_body["model"] = FALLBACK_MODEL
r = await client.post(fb_url, json=fallback_body, headers={"Authorization": f"Bearer {fb_key}"})
content_type = (r.headers.get("content-type") or "").lower()
if "application/json" in content_type:
try:
payload = r.json()
except Exception:
payload = {"error": "Upstream returned invalid JSON"}
else:
payload = {
"error": "Upstream returned non-JSON response",
"status_code": r.status_code,
"message": r.text[:1000],
}
return JSONResponse(status_code=r.status_code, content=payload)
# ──────────────────────────────────────────────
# PROMPT ANALYZE (/gen/prompt_analyze)
# ──────────────────────────────────────────────
@router.post("/prompt_analyze")
async def analyze_prompt(request: Request):
body = await request.json()
messages = body.get("prompt", [])
if not isinstance(messages, list) or len(messages) == 0:
raise HTTPException(400, "messages[] is required")
uses_tools = (
"tools" in body and isinstance(body["tools"], list) and len(body["tools"]) > 0
) or ("tool_choice" in body and body["tool_choice"] not in [None, "none"])
chosen_model, _ = route_chat(messages, uses_tools=uses_tools)
return {MODEL_MAP.get(chosen_model, chosen_model)}
# ──────────────────────────────────────────────
# MODELS LIST
# ──────────────────────────────────────────────
@router.get("/models")
def return_models_openai():
return {
"object": "list",
"data": [
{
"id": "lightning",
"object": "model",
"created": 1767225600,
"owned_by": "inferenceport-ai",
}
],
}
# ──────────────────────────────────────────────
# RESPONSES API (/gen/responses)
# ──────────────────────────────────────────────
def _resp_id(prefix: str) -> str:
return f"{prefix}_{uuid4().hex}"
def _resp_ts() -> int:
return int(time())
def _content_to_text(content: Any) -> str:
if isinstance(content, str):
return content
if isinstance(content, list):
parts = []
for item in content:
if isinstance(item, dict) and item.get("type") in ("input_text", "output_text", "text"):
txt = item.get("text")
if isinstance(txt, str):
parts.append(txt)
return "".join(parts)
return ""
def _responses_input_to_messages(
input_data: Any,
instructions: Optional[str] = None,
) -> List[Dict[str, Any]]:
messages: List[Dict[str, Any]] = []
if instructions:
messages.append({"role": "developer", "content": instructions})
if isinstance(input_data, str):
messages.append({"role": "user", "content": input_data})
return messages
if isinstance(input_data, list):
for item in input_data:
if isinstance(item, str):
messages.append({"role": "user", "content": item})
continue
if not isinstance(item, dict):
continue
role = item.get("role", "user")
text = _content_to_text(item.get("content", ""))
if text:
messages.append({"role": role, "content": text})
return messages
def _build_responses_payload(
model: str,
text: str,
response_id: str,
input_tokens: int = 0,
output_tokens: int = 0,
) -> Dict[str, Any]:
return {
"id": response_id,
"object": "response",
"created_at": _resp_ts(),
"status": "completed",
"completed_at": _resp_ts(),
"error": None,
"incomplete_details": None,
"instructions": None,
"max_output_tokens": None,
"model": model,
"output": [
{
"id": _resp_id("msg"),
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": text, "annotations": []}],
}
],
"output_text": text,
"usage": {
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
},
}
@router.post("/responses")
async def create_responses(
request: Request,
authorization: Optional[str] = Header(None),
x_client_id: Optional[str] = Header(None),
):
body = await request.json()
model = body.get("model")
input_data = body.get("input")
instructions = body.get("instructions")
stream = body.get("stream", True)
if not model:
raise HTTPException(400, "model is required")
if input_data is None:
raise HTTPException(400, "input is required")
messages = _responses_input_to_messages(input_data, instructions=instructions)
if not messages:
raise HTTPException(400, "input could not be parsed")
# ── shared helper: route + call + return (text, input_tokens, output_tokens) ──
async def _generate() -> Tuple[str, int, int]:
chosen_model, provider = route_chat(messages)
await _check_chat_rate_limit(request, authorization, x_client_id)
data = await call_chat_completions(messages, chosen_model, provider)
text = _extract_text_from_response(data)
input_tokens, output_tokens = _extract_usage(data)
return text, input_tokens, output_tokens
# ── non-streaming ─────────────────────────
if stream is False:
text, input_tokens, output_tokens = await _generate()
response_id = _resp_id("resp")
return JSONResponse(
content=_build_responses_payload(model, text, response_id, input_tokens, output_tokens)
)
# ── streaming ─────────────────────────────
async def event_stream():
response_id = _resp_id("resp")
created_evt = {
"type": "response.created",
"response": {
"id": response_id,
"object": "response",
"created_at": _resp_ts(),
"status": "in_progress",
"model": model,
},
}
yield f"data: {json.dumps(created_evt)}\n\n"
try:
text, input_tokens, output_tokens = await _generate()
except HTTPException as exc:
err_evt = {"type": "response.error", "error": {"message": exc.detail}}
yield f"data: {json.dumps(err_evt)}\n\n"
yield "data: [DONE]\n\n"
return
# Stream text in chunks
chunk_size = 64
for i in range(0, len(text), chunk_size):
delta_evt = {
"type": "response.output_text.delta",
"response_id": response_id,
"delta": text[i : i + chunk_size],
}
yield f"data: {json.dumps(delta_evt)}\n\n"
completed_evt = {
"type": "response.completed",
"response": _build_responses_payload(model, text, response_id, input_tokens, output_tokens),
}
yield f"data: {json.dumps(completed_evt)}\n\n"
yield "data: [DONE]\n\n"
return StreamingResponse(
event_stream(),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "Connection": "keep-alive", "X-Accel-Buffering": "no"},
)