Spaces:
Running
Running
| 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 | |
| # ────────────────────────────────────────────── | |
| 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 | |
| # ────────────────────────────────────────────── | |
| 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 | |
| # ────────────────────────────────────────────── | |
| 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) | |
| # ────────────────────────────────────────────── | |
| 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) | |
| # ────────────────────────────────────────────── | |
| 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) | |
| 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) | |
| # ────────────────────────────────────────────── | |
| 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 | |
| # ────────────────────────────────────────────── | |
| 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, | |
| }, | |
| } | |
| 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"}, | |
| ) |