9meo's picture
Add fine-tuned cross-border guardrail (OpenThai-SystemOne)
dc0dbce verified
Raw History Blame Contribute Delete
6.65 kB
# -*- coding: utf-8 -*-
"""
guardrail.py — production wrapper สำหรับ Pathumma Cross-border Guardrail
Pipeline: pre-filter (regex/DLP) → System One model → threshold → routing decision
- ใช้ instructions/criteria จาก guardrail_config.json เท่านั้น (ชุดเดียวกับที่เทรน)
- fail-closed: error / confidence ต่ำ / abstain สูง → sovereign_only + needs_review
- pre-filter จับ pattern ที่โมเดลไม่ได้เทรนให้จำ (secret จริง, เลขบัตร 13 หลัก) แล้ว override ขึ้นทันที
ใช้:
from guardrail import Guardrail
g = Guardrail("TODO-org/pathumma-crossborder-guardrail") # หรือ local path
g.check("ช่วยสรุปผลตรวจของ นายสมชาย ใจดี HN 123456")
"""
from __future__ import annotations
import json, os, re, time, logging
from dataclasses import dataclass, asdict
from typing import Dict, List, Optional
log = logging.getLogger("guardrail")
# ---------- deterministic pre-filters (แก้/เพิ่มได้ตาม policy องค์กร) ----------
_PREFILTERS: List[tuple[str, re.Pattern, str]] = [
# (category_key, pattern, reason)
("system_secret", re.compile(r"-----BEGIN (?:RSA |EC |OPENSSH )?PRIVATE KEY-----"), "private key block"),
("system_secret", re.compile(r"\b(?:sk|ghp|gho|xox[bap]|AKIA|AIza)[A-Za-z0-9_\-]{16,}"), "api-key-like token"),
("system_secret", re.compile(r"(?i)(?:password|passwd|pwd|secret|token)\s*[=:]\s*\S{6,}"), "credential assignment"),
("system_secret", re.compile(r"(?i)\b(?:postgres|mysql|mongodb|redis)://[^\s]+:[^\s]+@"), "connection string with credentials"),
("pdpa_general", re.compile(r"(?<!\d)\d-?\d{4}-?\d{5}-?\d{2}-?\d(?!\d)"), "thai national id pattern"),
]
def _thai_id_checksum_ok(s: str) -> bool:
d = re.sub(r"\D", "", s)
if len(d) != 13: return False
return (11 - sum(int(d[i]) * (13 - i) for i in range(12)) % 11) % 10 == int(d[12])
@dataclass
class Decision:
category: str
category_no: int
routing: str
confidence: float
abstain: float
needs_review: bool
source: str # "model" | "prefilter" | "fallback"
reason: str
probabilities: Dict[str, float]
latency_ms: float
def to_dict(self): return asdict(self)
class Guardrail:
def __init__(self, model_id_or_path: str, config_path: Optional[str] = None,
use_prefilter: bool = True, device: Optional[str] = None):
from openthai_systemone import SystemOneClient, Choice
self._Choice = Choice
self.cfg = self._load_config(model_id_or_path, config_path)
self.client = SystemOneClient(model_id_or_path)
self.use_prefilter = use_prefilter
c = self.cfg
self.question = {c["question_id"]: Choice(instructions=c["instructions"], criteria=c["criteria"])}
self.state_key = c.get("state_key", "prompt")
self.th = c["thresholds"]
self.fail_closed = c.get("fail_closed_routing", "sovereign_only")
@staticmethod
def _load_config(model_id_or_path: str, config_path: Optional[str]) -> dict:
if config_path and os.path.exists(config_path):
return json.load(open(config_path, encoding="utf-8"))
local = os.path.join(model_id_or_path, "guardrail_config.json")
if os.path.exists(local):
return json.load(open(local, encoding="utf-8"))
from huggingface_hub import hf_hub_download # ดึงจาก HF repo
p = hf_hub_download(model_id_or_path, "guardrail_config.json")
return json.load(open(p, encoding="utf-8"))
# ---------- public API ----------
def check(self, text: str) -> Decision:
t0 = time.perf_counter()
c = self.cfg
# 1) deterministic pre-filter
if self.use_prefilter:
hit = self._prefilter(text)
if hit:
cat, reason = hit
return Decision(cat, c["category_no"][cat], c["routing"][cat], 1.0, 0.0, False,
"prefilter", reason, {cat: 1.0}, (time.perf_counter() - t0) * 1000)
# 2) model
try:
resp = self.client.system_one(state={self.state_key: text}, questions=self.question)
a = resp.answers[c["question_id"]]
except Exception as e: # fail-closed
log.exception("guardrail model error")
return Decision("unknown", -1, self.fail_closed, 0.0, 1.0, True, "fallback",
f"model error: {type(e).__name__}", {}, (time.perf_counter() - t0) * 1000)
needs_review = (a.confidence < self.th["min_confidence"]) or (a.abstain > self.th["max_abstain"])
routing = self.fail_closed if needs_review else c["routing"][a.choice]
reason = "low confidence/abstain" if needs_review else "model decision"
return Decision(a.choice, c["category_no"][a.choice], routing, float(a.confidence), float(a.abstain),
needs_review, "model", reason, dict(a.probabilities), (time.perf_counter() - t0) * 1000)
def check_batch(self, texts: List[str]) -> List[Decision]:
return [self.check(t) for t in texts] # client ยังไม่มี batch API; วน loop (≈40-150 ms/req)
# ---------- helpers ----------
def _prefilter(self, text: str):
for cat, pat, reason in _PREFILTERS:
m = pat.search(text)
if not m: continue
if reason == "thai national id pattern" and not _thai_id_checksum_ok(m.group(0)):
continue # เลขสุ่มที่ checksum ไม่ผ่าน ปล่อยให้โมเดลตัดสิน
return cat, reason
return None
if __name__ == "__main__":
import sys
g = Guardrail(sys.argv[1] if len(sys.argv) > 1 else "./systemone-guardrail-ft")
for s in ["อธิบาย transformer ให้เด็กเข้าใจ",
"ช่วยสรุปผลตรวจสุขภาพของ นายสมชาย ใจดี HN 123456 พบเบาหวาน",
"ช่วยดู .env นี้ DB_PASSWORD=Pa55w0rd123 DB_HOST=10.0.0.5",
"ช่วยแปลรายงานข่าวกรองชั้นลับมากของ สมช."]:
print(json.dumps(g.check(s).to_dict(), ensure_ascii=False, indent=1))