feat: classify sensitive columns locally
This commit is contained in:
@@ -0,0 +1,301 @@
|
||||
"""Offline, CPU-only JSONL worker for optional sensitivity NER evidence."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import ctypes
|
||||
import errno
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
PII_LABELS = [
|
||||
"person",
|
||||
"full_name",
|
||||
"first_name",
|
||||
"middle_name",
|
||||
"last_name",
|
||||
"date_of_birth",
|
||||
"email",
|
||||
"phone_number",
|
||||
"address",
|
||||
"street_address",
|
||||
"city",
|
||||
"state_or_region",
|
||||
"postal_code",
|
||||
"country",
|
||||
"government_id",
|
||||
"national_id_number",
|
||||
"passport_number",
|
||||
"drivers_license_number",
|
||||
"license_number",
|
||||
"tax_id",
|
||||
"tax_number",
|
||||
"bank_account",
|
||||
"account_number",
|
||||
"routing_number",
|
||||
"iban",
|
||||
"payment_card",
|
||||
"card_number",
|
||||
"card_expiry",
|
||||
"card_cvv",
|
||||
"username",
|
||||
"ip_address",
|
||||
"account_id",
|
||||
"sensitive_account_id",
|
||||
"password",
|
||||
"secret",
|
||||
"api_key",
|
||||
"access_token",
|
||||
"recovery_code",
|
||||
"sensitive_date",
|
||||
"document_date",
|
||||
"expiration_date",
|
||||
"transaction_date",
|
||||
]
|
||||
|
||||
_MODEL_COMPAT_DIRECTORY: tempfile.TemporaryDirectory[str] | None = None
|
||||
_EXPECTED_MODEL_REVISION = "c153999da5f4c509df4322b0c6a1baf3d2c284d7"
|
||||
|
||||
|
||||
def _arguments() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(add_help=False)
|
||||
parser.add_argument("--model", required=True)
|
||||
parser.add_argument("--threads", type=int, default=2)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _disable_network() -> None:
|
||||
libc = ctypes.CDLL(None, use_errno=True)
|
||||
libc.prctl.argtypes = [
|
||||
ctypes.c_int,
|
||||
ctypes.c_ulong,
|
||||
ctypes.c_ulong,
|
||||
ctypes.c_ulong,
|
||||
ctypes.c_ulong,
|
||||
]
|
||||
libc.prctl.restype = ctypes.c_int
|
||||
if libc.prctl(38, 1, 0, 0, 0) != 0: # PR_SET_NO_NEW_PRIVS
|
||||
raise RuntimeError("cannot enable no-new-privileges for network isolation")
|
||||
|
||||
try:
|
||||
seccomp = ctypes.CDLL("libseccomp.so.2", use_errno=True)
|
||||
except OSError as error:
|
||||
raise RuntimeError("libseccomp is required for network isolation") from error
|
||||
seccomp.seccomp_init.argtypes = [ctypes.c_uint32]
|
||||
seccomp.seccomp_init.restype = ctypes.c_void_p
|
||||
seccomp.seccomp_syscall_resolve_name.argtypes = [ctypes.c_char_p]
|
||||
seccomp.seccomp_syscall_resolve_name.restype = ctypes.c_int
|
||||
seccomp.seccomp_rule_add.argtypes = [
|
||||
ctypes.c_void_p,
|
||||
ctypes.c_uint32,
|
||||
ctypes.c_int,
|
||||
ctypes.c_uint,
|
||||
]
|
||||
seccomp.seccomp_rule_add.restype = ctypes.c_int
|
||||
seccomp.seccomp_load.argtypes = [ctypes.c_void_p]
|
||||
seccomp.seccomp_load.restype = ctypes.c_int
|
||||
seccomp.seccomp_release.argtypes = [ctypes.c_void_p]
|
||||
seccomp.seccomp_release.restype = None
|
||||
|
||||
allow = 0x7FFF0000 # SCMP_ACT_ALLOW
|
||||
deny = 0x00050000 | errno.EPERM # SCMP_ACT_ERRNO(EPERM)
|
||||
filter_context = seccomp.seccomp_init(allow)
|
||||
if not filter_context:
|
||||
raise RuntimeError("cannot initialize network syscall filter")
|
||||
try:
|
||||
for syscall in (
|
||||
"socket",
|
||||
"connect",
|
||||
"sendto",
|
||||
"sendmsg",
|
||||
"sendmmsg",
|
||||
"bind",
|
||||
"listen",
|
||||
"accept",
|
||||
"accept4",
|
||||
):
|
||||
syscall_number = seccomp.seccomp_syscall_resolve_name(syscall.encode("ascii"))
|
||||
if syscall_number < 0:
|
||||
raise RuntimeError(f"cannot resolve network syscall: {syscall}")
|
||||
if seccomp.seccomp_rule_add(filter_context, deny, syscall_number, 0) != 0:
|
||||
raise RuntimeError(f"cannot block network syscall: {syscall}")
|
||||
if seccomp.seccomp_load(filter_context) != 0:
|
||||
raise RuntimeError("cannot activate network syscall filter")
|
||||
finally:
|
||||
seccomp.seccomp_release(filter_context)
|
||||
|
||||
def blocked(*_args: Any, **_kwargs: Any) -> Any:
|
||||
raise PermissionError(errno.EPERM, "network disabled")
|
||||
|
||||
socket.socket = blocked # type: ignore[assignment]
|
||||
socket.create_connection = blocked # type: ignore[assignment]
|
||||
|
||||
|
||||
def _verify_model(path: Path) -> None:
|
||||
revision_path = path / "THOTHII_MODEL_REVISION"
|
||||
try:
|
||||
revision = revision_path.read_text(encoding="utf-8").strip()
|
||||
except OSError as error:
|
||||
raise RuntimeError("model revision marker is unavailable") from error
|
||||
if revision != _EXPECTED_MODEL_REVISION:
|
||||
raise RuntimeError("model revision is not approved")
|
||||
|
||||
manifest_path = Path(__file__).with_name("sensitivity-ner-model-sha256.txt")
|
||||
try:
|
||||
manifest = manifest_path.read_text(encoding="utf-8").splitlines()
|
||||
except OSError as error:
|
||||
raise RuntimeError("model checksum manifest is unavailable") from error
|
||||
for line in manifest:
|
||||
checksum, separator, relative_name = line.partition(" ")
|
||||
if not separator or len(checksum) != 64 or not relative_name.startswith("./"):
|
||||
raise RuntimeError("model checksum manifest is invalid")
|
||||
relative_path = Path(relative_name[2:])
|
||||
if relative_path.is_absolute() or ".." in relative_path.parts:
|
||||
raise RuntimeError("model checksum path is invalid")
|
||||
model_file = path / relative_path
|
||||
if not model_file.is_file() or model_file.is_symlink():
|
||||
raise RuntimeError("approved model file is unavailable")
|
||||
digest = hashlib.sha256()
|
||||
with model_file.open("rb") as stream:
|
||||
for chunk in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
if digest.hexdigest() != checksum:
|
||||
raise RuntimeError("approved model checksum does not match")
|
||||
|
||||
|
||||
def _transformers4_model_path(path: Path) -> Path:
|
||||
"""Adapt tokenizer metadata emitted by Transformers 5 without changing pinned weights.
|
||||
|
||||
GLiNER2 2.0.0 officially requires Transformers <5, while current Fastino checkpoints were
|
||||
saved by Transformers 5.8.0. Transformers 4 calls the same list
|
||||
``additional_special_tokens``; Transformers 5 renamed it to ``extra_special_tokens`` and
|
||||
changed its type. Keep the downloaded model immutable and create a temporary symlink view
|
||||
containing only the compatibility metadata needed by the supported GLiNER2 dependency set.
|
||||
"""
|
||||
|
||||
tokenizer_path = path / "tokenizer_config.json"
|
||||
try:
|
||||
tokenizer = json.loads(tokenizer_path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError) as error:
|
||||
raise RuntimeError("invalid tokenizer configuration") from error
|
||||
extra_tokens = tokenizer.get("extra_special_tokens")
|
||||
if extra_tokens is None:
|
||||
return path
|
||||
if not isinstance(extra_tokens, list) or not all(isinstance(token, str) for token in extra_tokens):
|
||||
raise RuntimeError("unsupported extra_special_tokens configuration")
|
||||
if "additional_special_tokens" in tokenizer:
|
||||
raise RuntimeError("ambiguous special-token configuration")
|
||||
|
||||
global _MODEL_COMPAT_DIRECTORY
|
||||
_MODEL_COMPAT_DIRECTORY = tempfile.TemporaryDirectory(prefix="thothii-ner-model-")
|
||||
compatible_path = Path(_MODEL_COMPAT_DIRECTORY.name)
|
||||
for child in path.iterdir():
|
||||
if child.name == tokenizer_path.name:
|
||||
continue
|
||||
(compatible_path / child.name).symlink_to(child, target_is_directory=child.is_dir())
|
||||
tokenizer["additional_special_tokens"] = tokenizer.pop("extra_special_tokens")
|
||||
(compatible_path / tokenizer_path.name).write_text(
|
||||
json.dumps(tokenizer, ensure_ascii=False, indent=2) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
return compatible_path
|
||||
|
||||
|
||||
def _load_model(model_path: str, threads: int) -> Any:
|
||||
path = Path(model_path).resolve(strict=True)
|
||||
if not path.is_dir():
|
||||
raise RuntimeError("model path must be a local directory")
|
||||
_verify_model(path)
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
||||
os.environ["HIP_VISIBLE_DEVICES"] = ""
|
||||
os.environ["HF_HUB_OFFLINE"] = "1"
|
||||
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
||||
import torch
|
||||
from gliner2 import AutoExtractor
|
||||
|
||||
torch.set_num_threads(max(1, min(threads, 8)))
|
||||
torch.set_num_interop_threads(1)
|
||||
compatible_path = _transformers4_model_path(path)
|
||||
with contextlib.redirect_stdout(sys.stderr):
|
||||
model = AutoExtractor.from_pretrained(str(compatible_path), map_location="cpu")
|
||||
_disable_network()
|
||||
return model
|
||||
|
||||
|
||||
def _request(value: Any) -> tuple[str, list[dict[str, str]]]:
|
||||
if not isinstance(value, dict) or not isinstance(value.get("id"), str):
|
||||
raise ValueError("invalid request")
|
||||
candidates = value.get("candidates")
|
||||
if not isinstance(candidates, list) or not 1 <= len(candidates) <= 128:
|
||||
raise ValueError("invalid candidates")
|
||||
parsed: list[dict[str, str]] = []
|
||||
for candidate in candidates:
|
||||
if not isinstance(candidate, dict):
|
||||
raise ValueError("invalid candidate")
|
||||
column_id = candidate.get("columnId")
|
||||
text = candidate.get("text")
|
||||
if not isinstance(column_id, str) or not isinstance(text, str) or not 1 <= len(text) <= 500:
|
||||
raise ValueError("invalid candidate")
|
||||
parsed.append({"columnId": column_id, "text": text})
|
||||
return value["id"], parsed
|
||||
|
||||
|
||||
def _detect(model: Any, candidates: list[dict[str, str]]) -> list[dict[str, Any]]:
|
||||
evidence: list[dict[str, Any]] = []
|
||||
for candidate in candidates:
|
||||
result = model.extract_entities(
|
||||
candidate["text"],
|
||||
PII_LABELS,
|
||||
threshold=0.5,
|
||||
include_confidence=True,
|
||||
)
|
||||
entities = result.get("entities", {}) if isinstance(result, dict) else {}
|
||||
best: tuple[str, float] | None = None
|
||||
if isinstance(entities, dict):
|
||||
for label, matches in entities.items():
|
||||
if label not in PII_LABELS or not isinstance(matches, list):
|
||||
continue
|
||||
for match in matches:
|
||||
if not isinstance(match, dict):
|
||||
continue
|
||||
confidence = match.get("confidence")
|
||||
if not isinstance(confidence, (int, float)) or not 0 <= confidence <= 1:
|
||||
continue
|
||||
if best is None or confidence > best[1]:
|
||||
best = (label, float(confidence))
|
||||
if best is not None:
|
||||
evidence.append(
|
||||
{
|
||||
"columnId": candidate["columnId"],
|
||||
"label": best[0],
|
||||
"confidence": best[1],
|
||||
}
|
||||
)
|
||||
return evidence
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = _arguments()
|
||||
model = _load_model(args.model, args.threads)
|
||||
print(json.dumps({"ready": True}, separators=(",", ":")), flush=True)
|
||||
for line in sys.stdin:
|
||||
request_id = "invalid"
|
||||
try:
|
||||
request_id, candidates = _request(json.loads(line))
|
||||
response = {"id": request_id, "ok": True, "evidence": _detect(model, candidates)}
|
||||
except Exception:
|
||||
response = {"id": request_id, "ok": False, "error": "detection_failed"}
|
||||
print(json.dumps(response, separators=(",", ":")), flush=True)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user