Spaces:
Sleeping
Sleeping
| import os | |
| import time | |
| from typing import List, Optional, Any | |
| from uuid import uuid4 | |
| import google.generativeai as genai | |
| from fastapi import FastAPI, HTTPException | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from pydantic import BaseModel, Field | |
| API_KEY = os.getenv("GEMINI_API_KEY", "").strip() | |
| if not API_KEY: | |
| raise RuntimeError("Defina GEMINI_API_KEY no Hugging Face Space.") | |
| genai.configure(api_key=API_KEY) | |
| _model = None | |
| _model_name = None | |
| def _clean_model_name(name: str) -> str: | |
| name = (name or "").strip() | |
| return name[7:] if name.startswith("models/") else name | |
| def _available_generate_models() -> List[str]: | |
| names = [] | |
| try: | |
| for item in genai.list_models(): | |
| methods = getattr(item, "supported_generation_methods", None) or [] | |
| if "generateContent" in methods: | |
| name = _clean_model_name(getattr(item, "name", "")) | |
| if name: | |
| names.append(name) | |
| except Exception: | |
| pass | |
| return names | |
| def get_model(): | |
| global _model, _model_name | |
| if _model is not None: | |
| return _model | |
| requested = _clean_model_name(os.getenv("GEMINI_MODEL", "")) | |
| available = _available_generate_models() | |
| if requested: | |
| chosen = requested | |
| else: | |
| preferred = [ | |
| "gemini-2.5-flash", | |
| "gemini-2.0-flash", | |
| "gemini-flash-latest", | |
| "gemini-pro-latest", | |
| ] | |
| chosen = next((x for x in preferred if x in available), None) | |
| if not chosen and available: | |
| chosen = available[0] | |
| if not chosen: | |
| chosen = _clean_model_name( | |
| os.getenv("GEMINI_FALLBACK_MODEL", "gemini-2.0-flash") | |
| ) | |
| _model_name = chosen | |
| _model = genai.GenerativeModel(chosen) | |
| return _model | |
| def current_model_name() -> str: | |
| return _model_name or _clean_model_name(os.getenv("GEMINI_MODEL", "")) or "auto" | |
| app = FastAPI(title="Gemini OpenAI-Compatible Endpoint", version="2.0.0") | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], | |
| allow_credentials=False, | |
| allow_methods=["*"], | |
| allow_headers=["*"], | |
| ) | |
| class PromptRequest(BaseModel): | |
| prompt: str | |
| max_new_tokens: int = Field(default=300, ge=1, le=8192) | |
| temperature: float = Field(default=0.7, ge=0.0, le=2.0) | |
| history: List[str] = Field(default_factory=list) | |
| class GenerateSoalRequest(BaseModel): | |
| topic: str | |
| level: int = Field(default=1, ge=1, le=5) | |
| tipe_pertanyaan: Optional[str] = None | |
| class ChatMessage(BaseModel): | |
| role: str | |
| content: Any | |
| class ChatCompletionRequest(BaseModel): | |
| model: Optional[str] = None | |
| messages: List[ChatMessage] | |
| temperature: float = Field(default=0.7, ge=0.0, le=2.0) | |
| max_tokens: int = Field(default=512, ge=1, le=8192) | |
| stream: bool = False | |
| def normalize_content(content: Any) -> str: | |
| if isinstance(content, str): | |
| return content | |
| if isinstance(content, list): | |
| parts = [] | |
| for item in content: | |
| if isinstance(item, dict) and "text" in item: | |
| parts.append(str(item["text"])) | |
| else: | |
| parts.append(str(item)) | |
| return "\n".join(parts) | |
| return "" if content is None else str(content) | |
| def chat_messages_to_gemini(messages: List[ChatMessage]): | |
| result = [] | |
| system = [] | |
| for msg in messages: | |
| role = (msg.role or "user").lower() | |
| text = normalize_content(msg.content) | |
| if role == "system": | |
| system.append(text) | |
| continue | |
| result.append({ | |
| "role": "model" if role == "assistant" else "user", | |
| "parts": [{"text": text}], | |
| }) | |
| if system: | |
| prefix = "INSTRUÇÕES DO SISTEMA:\n" + "\n\n".join(system) + "\n\n" | |
| if result and result[0]["role"] == "user": | |
| result[0]["parts"][0]["text"] = prefix + result[0]["parts"][0]["text"] | |
| else: | |
| result.insert(0, {"role": "user", "parts": [{"text": prefix}]}) | |
| return result or [{"role": "user", "parts": [{"text": "Olá"}]}] | |
| def extract_text(result) -> str: | |
| try: | |
| return result.text | |
| except Exception: | |
| return str(result) | |
| def format_prompt_soal(topic, tipe=None, level=1): | |
| tipe_instrucao = f"Tipe soal solicitado: {tipe}." if tipe else "Tipo de questão livre." | |
| return f""" | |
| Crie uma questão de prática de gramática inglesa sobre: "{topic}". | |
| {tipe_instrucao} | |
| Nível: {level}/5. | |
| Responda SOMENTE com JSON válido com as chaves: | |
| text_pertanyaan, tipe_pertanyaan, opsi, jawaban_benar, penjelasan. | |
| """.strip() | |
| async def root(): | |
| return { | |
| "status": "online", | |
| "service": "Gemini endpoint", | |
| "model": current_model_name(), | |
| "endpoints": { | |
| "generate": "POST /generate", | |
| "chat": "POST /v1/chat/completions", | |
| "models": "GET /v1/models", | |
| "health": "GET /health", | |
| }, | |
| } | |
| async def root_head(): | |
| return {} | |
| async def health(): | |
| return {"status": "ok", "model": current_model_name()} | |
| async def models(): | |
| available = _available_generate_models() | |
| selected = current_model_name() | |
| if selected == "auto" and available: | |
| selected = available[0] | |
| return { | |
| "object": "list", | |
| "data": [{ | |
| "id": selected, | |
| "object": "model", | |
| "created": int(time.time()), | |
| "owned_by": "google", | |
| }], | |
| } | |
| async def generate_text(req: PromptRequest): | |
| try: | |
| messages = [] | |
| for msg in req.history: | |
| if msg.startswith("Bot:") or msg.startswith("Assistant:"): | |
| role = "model" | |
| text = msg.split(":", 1)[1].strip() if ":" in msg else msg | |
| else: | |
| role = "user" | |
| text = msg.split(":", 1)[1].strip() if msg.startswith("User:") else msg | |
| messages.append({"role": role, "parts": [{"text": text}]}) | |
| messages.append({"role": "user", "parts": [{"text": req.prompt}]}) | |
| result = get_model().generate_content( | |
| messages, | |
| generation_config={ | |
| "temperature": req.temperature, | |
| "max_output_tokens": req.max_new_tokens, | |
| }, | |
| ) | |
| return {"generated_text": extract_text(result), "model": current_model_name()} | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=str(e)) | |
| async def chat_completions(req: ChatCompletionRequest): | |
| if req.stream: | |
| raise HTTPException(status_code=400, detail="stream=true ainda não é suportado.") | |
| try: | |
| result = get_model().generate_content( | |
| chat_messages_to_gemini(req.messages), | |
| generation_config={ | |
| "temperature": req.temperature, | |
| "max_output_tokens": req.max_tokens, | |
| }, | |
| ) | |
| answer = extract_text(result) | |
| return { | |
| "id": "chatcmpl-" + uuid4().hex, | |
| "object": "chat.completion", | |
| "created": int(time.time()), | |
| "model": current_model_name(), | |
| "choices": [{ | |
| "index": 0, | |
| "message": {"role": "assistant", "content": answer}, | |
| "finish_reason": "stop", | |
| }], | |
| } | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=str(e)) | |
| async def generate_soal(req: GenerateSoalRequest): | |
| try: | |
| import json | |
| prompt = format_prompt_soal(req.topic, req.tipe_pertanyaan, req.level) | |
| result = get_model().generate_content( | |
| prompt, | |
| generation_config={"temperature": 0.8, "max_output_tokens": 512}, | |
| ) | |
| text = extract_text(result).strip() | |
| if text.startswith("```"): | |
| lines = text.splitlines()[1:] | |
| if lines and lines[-1].strip().startswith("```"): | |
| lines = lines[:-1] | |
| text = "\n".join(lines).strip() | |
| if text.lower().startswith("json"): | |
| text = text[4:].lstrip() | |
| soal_data = json.loads(text) | |
| soal_data["id"] = str(uuid4()) | |
| soal_data["topic_id"] = req.topic | |
| soal_data["level"] = req.level | |
| soal_data["is_ai_generated"] = True | |
| return {"soal": soal_data} | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Gagal generate soal: {e}") |