from fastapi import FastAPI, Request, HTTPException, Depends from fastapi.security import APIKeyHeader from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse, HTMLResponse from fastapi.openapi.utils import get_openapi import os import json import time import uuid from llama_cpp import Llama app = FastAPI(title="KAI API") app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"], ) def custom_openapi(): if app.openapi_schema: return app.openapi_schema schema = get_openapi( title="KAI API", version="0.1.0", routes=app.routes, ) schema["components"]["securitySchemes"] = { "ApiKeyAuth": { "type": "apiKey", "in": "header", "name": "X-API-Key" } } schema["security"] = [{"ApiKeyAuth": []}] app.openapi_schema = schema return schema app.openapi = custom_openapi API_KEY = os.environ.get("QWEN_API_KEY") MODEL_PATH = "/models/Qwen2.5-7B-Instruct-Q4_K_M.gguf" print(f"Loading model from {MODEL_PATH}...") llm = Llama( model_path=MODEL_PATH, n_ctx=4096, n_threads=8, n_gpu_layers=-1, verbose=False ) print("Model loaded successfully.") api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False) async def verify_key(api_key: str = Depends(api_key_header)): if not API_KEY: raise HTTPException(status_code=500, detail="QWEN_API_KEY not set in HF Secrets") if api_key != API_KEY: raise HTTPException(status_code=401, detail="Invalid API key") return api_key @app.get("/", response_class=HTMLResponse) async def ui(): return """
Odpowiedź pojawi się tutaj...""" @app.get("/health") async def health(): return { "status": "ok", "model": MODEL_PATH, "api_key_loaded": API_KEY is not None } def format_prompt(messages): prompt = "" for msg in messages: role = msg["role"] content = msg["content"] prompt += f"<|im_start|>{role}\n{content}<|im_end|>\n" prompt += "<|im_start|>assistant\n" return prompt @app.post("/v1/chat/completions") async def chat_completions(request: Request, api_key: str = Depends(verify_key)): body = await request.json() messages = body.get("messages", []) max_tokens = body.get("max_tokens", 512) temperature = body.get("temperature", 0.7) top_p = body.get("top_p", 0.95) if not messages: raise HTTPException(status_code=400, detail="messages is required") prompt = format_prompt(messages) output = llm( prompt, max_tokens=max_tokens, temperature=temperature, top_p=top_p, stop=["<|im_end|>"] ) response_text = output["choices"][0]["text"].strip() return JSONResponse({ "id": f"chatcmpl-{uuid.uuid4()}", "object": "chat.completion", "created": int(time.time()), "model": MODEL_PATH, "choices": [{ "index": 0, "message": {"role": "assistant", "content": response_text}, "finish_reason": output["choices"][0]["finish_reason"] }], "usage": { "prompt_tokens": output["usage"]["prompt_tokens"], "completion_tokens": output["usage"]["completion_tokens"], "total_tokens": output["usage"]["total_tokens"] } })