guard-a40M / server.py
cortiqa112's picture
Release Cortiqa Guard-A40M: Universal AI Security Firewall (v2 Enterprise)
ad966b1 verified
Raw History Blame Contribute Delete
3.16 kB
"""
Cortiqa-Guard-40M Production HTTP Microservice (FastAPI)
Deploy anywhere: Docker, Cloud (AWS/GCP/Azure), VPS, Kubernetes, or Local Server.
"""
import time
from pathlib import Path
from typing import Dict, List, Optional
import torch
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from transformers import AutoTokenizer
from model import CortiqaGuard40M, load_config, GUARD_LABELS
app = FastAPI(
title="Cortiqa-Guard-40M Universal AI Firewall",
description="High-Speed (<3ms) Prompt Injection & Safety Guardrail API",
version="1.0.0",
)
# Global model container
guard_model: Optional[CortiqaGuard40M] = None
guard_tokenizer = None
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
class ScanRequest(BaseModel):
prompt: str
max_length: Optional[int] = 512
class ScanResponse(BaseModel):
is_safe: bool
action: str # "ALLOW" or "BLOCK"
threat_category: str
class_id: int
confidence: float
latency_ms: float
scores: Dict[str, float]
@app.on_event("startup")
def load_firewall():
global guard_model, guard_tokenizer
cfg_path = Path(__file__).parent / "configs" / "guard_40m.yaml"
cfg = load_config(str(cfg_path))
print(f"Loading tokenizer: {cfg.training.tokenizer_name}...")
guard_tokenizer = AutoTokenizer.from_pretrained(cfg.training.tokenizer_name)
if guard_tokenizer.pad_token is None:
guard_tokenizer.pad_token = guard_tokenizer.eos_token
print(f"Loading Cortiqa Guard 40M on {device}...")
guard_model = CortiqaGuard40M(cfg.model).to(device)
# Check for trained checkpoint
ckpt_path = Path(__file__).parent / "checkpoints" / "guard_final.pt"
if ckpt_path.exists():
print(f"Loading trained weights from {ckpt_path}...")
ckpt = torch.load(ckpt_path, map_location=device)
guard_model.load_state_dict(ckpt["model_state"])
guard_model.eval()
print("✅ Cortiqa Guard Firewall API is live and ready to protect LLMs!")
@app.get("/health")
def health():
return {
"status": "healthy",
"device": str(device),
"parameters": guard_model.param_count() if guard_model else 0,
"firewall_active": True,
}
@app.post("/v1/scan", response_model=ScanResponse)
def scan_prompt(req: ScanRequest):
if not req.prompt.strip():
raise HTTPException(status_code=400, detail="Empty prompt provided.")
t0 = time.perf_counter()
enc = guard_tokenizer(
req.prompt,
return_tensors="pt",
max_length=req.max_length,
truncation=True,
).to(device)
with torch.no_grad():
res = guard_model.inspect_prompt(enc["input_ids"])
latency = (time.perf_counter() - t0) * 1000.0
return ScanResponse(
is_safe=res["is_safe"],
action="ALLOW" if res["is_safe"] else "BLOCK",
threat_category=res["label"],
class_id=res["class_id"],
confidence=res["confidence"],
latency_ms=round(latency, 2),
scores=res["all_scores"],
)
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)