import argparse import bisect import json import re import sys import unicodedata from dataclasses import asdict, dataclass from pathlib import Path from typing import Callable, List, Optional, Sequence, Tuple _HERE = Path(__file__).resolve().parent if str(_HERE) not in sys.path: # чтобы `core/` нашёлся из любого cwd sys.path.insert(0, str(_HERE)) WINDOW_WORDPIECES = 384 DEFAULT_RULES = _HERE / "gitleaks.toml" _BOUNDARY_RE = re.compile(r"(?:\n+|[.!?]\s+)") @dataclass(frozen=True) class Span: start: int end: int label: str source: str # ner | regex | both score: Optional[float] # None у спанов детерминированного слоя text: str def plan_windows( text: str, count: Callable[[str], int], maximum: int = WINDOW_WORDPIECES ) -> List[Tuple[int, int]]: """Режет текст на куски не длиннее `maximum` wordpiece, по границам предложений и абзацев там, где это возможно. Возвращает пары (начало, конец) в символах.""" if not text: return [] if count(text) <= maximum: return [(0, len(text))] boundaries = sorted({0, *(m.end() for m in _BOUNDARY_RE.finditer(text)), len(text)}) def fitting_end(left: int) -> int: low, step = left, maximum * 4 high = min(len(text), left + step) while high < len(text) and count(text[left:high]) <= maximum: low, step = high, step * 2 high = min(len(text), left + step) if high == len(text) and count(text[left:high]) <= maximum: return high while low + 1 < high: middle = (low + high) // 2 if count(text[left:middle]) <= maximum: low = middle else: high = middle return low windows: List[Tuple[int, int]] = [] start = 0 while start < len(text): candidate = fitting_end(start) if candidate <= start: raise RuntimeError(f"ни один непустой префикс не влезает в {maximum} wordpiece на {start}") boundary = boundaries[bisect.bisect_right(boundaries, candidate) - 1] end = boundary if boundary > start else candidate if count(text[start:end]) > maximum: end = candidate windows.append((start, end)) if end >= len(text): break start = end return windows def extract_entities(bio: Sequence[str]) -> List[Tuple[int, int, str]]: """BIO-теги → (начало, конец, метка) в индексах токенов. I- без своего B- открывает сущность: модель не обязана быть согласованной, а терять предсказание из-за этого нельзя.""" ents: List[Tuple[int, int, str]] = [] cur_label: Optional[str] = None cur_start = 0 for i, tag in enumerate(list(bio) + ["O"]): if tag.startswith("B-"): if cur_label is not None: ents.append((cur_start, i, cur_label)) cur_label, cur_start = tag[2:], i elif tag.startswith("I-"): label = tag[2:] if cur_label == label: continue if cur_label is not None: ents.append((cur_start, i, cur_label)) cur_label, cur_start = label, i else: if cur_label is not None: ents.append((cur_start, i, cur_label)) cur_label = None return ents def _lower_score(left: Optional[float], right: Optional[float]) -> Optional[float]: present = [s for s in (left, right) if s is not None] return min(present) if present else None def merge(spans: List[Span], text: str) -> List[Span]: """Склеивает пересекающиеся и примыкающие однометочные спаны. Метку задаёт первый спан — самый левый, затем самый длинный, затем NER вперёд regex.""" if not spans: return [] ordered = sorted( spans, key=lambda s: (s.start, -(s.end - s.start), 0 if s.source == "ner" else 1), ) out: List[Span] = [ordered[0]] for sp in ordered[1:]: prev = out[-1] if sp.start < prev.end or (sp.start == prev.end and sp.label == prev.label): end = max(prev.end, sp.end) out[-1] = Span( prev.start, end, prev.label, prev.source if sp.source == prev.source else "both", prev.score if sp.start < prev.end else _lower_score(prev.score, sp.score), text[prev.start:end], ) else: out.append(sp) return out class Detector: def __init__( self, model_id: str = "fef2/ner_rus_bert-secret_detection", device: str = "cpu", fp16: Optional[bool] = None, window_wordpieces: int = WINDOW_WORDPIECES, batch_size: int = 32, rules: Optional[str | Path] = DEFAULT_RULES, ) -> None: import torch from transformers import AutoModelForTokenClassification, AutoTokenizer self.torch = torch self.device = device self.window = window_wordpieces self.max_length = window_wordpieces + 2 # [CLS] … [SEP] self.batch_size = max(1, batch_size) self.tokenizer = AutoTokenizer.from_pretrained(model_id, use_fast=True) if not self.tokenizer.is_fast: raise SystemExit("нужен fast-токенизатор: без offset_mapping смещения не восстановить") model = AutoModelForTokenClassification.from_pretrained(model_id).eval() if fp16 is None: fp16 = device.startswith("cuda") if fp16: model = model.half() self.model = model.to(device) self.id2label = {int(k): v for k, v in model.config.id2label.items()} self.rules: list = [] if rules is not None: from core.scrubber import load_gitleaks_rules self.rules, skipped = load_gitleaks_rules(Path(rules)) if not self.rules: raise SystemExit(f"правила не загрузились: {rules}") self.rules_skipped = skipped def _count(self, text: str) -> int: return len(self.tokenizer(text, add_special_tokens=False)["input_ids"]) def spans(self, text: str) -> List[Span]: """Спаны по NFC-нормализованному тексту. Нормализуйте вход тем же NFC, прежде чем резать его по этим индексам.""" text = unicodedata.normalize("NFC", text) found: List[Span] = [] windows = plan_windows(text, self._count, self.window) for i in range(0, len(windows), self.batch_size): found.extend(self._forward(text, windows[i:i + self.batch_size])) found.extend(self.regex_spans(text)) return merge(found, text) def regex_spans(self, text: str) -> List[Span]: """Детерминированный слой: kv-детектор, CLI-детектор и правила gitleaks по непрозрачным значениям. Без модели, без GPU.""" if not self.rules: return [] from core.scrubber import credential_sites return [Span(s.start, s.end, s.label, "regex", None, text[s.start:s.end]) for s in credential_sites(text, self.rules)] def _forward(self, text: str, windows: Sequence[Tuple[int, int]]) -> List[Span]: torch = self.torch enc = self.tokenizer( [text[a:b] for a, b in windows], return_offsets_mapping=True, return_special_tokens_mask=True, truncation=True, max_length=self.max_length, padding=True, return_tensors="pt", ) with torch.inference_mode(): logits = self.model( input_ids=enc["input_ids"].to(self.device), attention_mask=enc["attention_mask"].to(self.device), ).logits preds = logits.argmax(-1).cpu() confidence = logits.float().softmax(-1).max(-1).values.cpu() offsets = enc["offset_mapping"].tolist() special = enc["special_tokens_mask"].tolist() out: List[Span] = [] for bi, (base, _) in enumerate(windows): real = [i for i, m in enumerate(special[bi]) if m == 0] bio = [self.id2label[int(preds[bi, i])] for i in real] for begin, end, label in extract_entities(bio): start = base + offsets[bi][real[begin]][0] stop = base + offsets[bi][real[end - 1]][1] if stop > start: score = min(float(confidence[bi, real[i]]) for i in range(begin, end)) out.append(Span(start, stop, label, "ner", round(score, 4), text[start:stop])) return out def mask(self, text: str, template: str = "[REDACTED:{label}]") -> str: text = unicodedata.normalize("NFC", text) parts, cursor = [], 0 for sp in self.spans(text): parts.append(text[cursor:sp.start]) parts.append(template.format(label=sp.label)) cursor = sp.end parts.append(text[cursor:]) return "".join(parts) def main() -> int: ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) src = ap.add_mutually_exclusive_group() src.add_argument("--text", help="текст прямо в аргументе") src.add_argument("--file", help="файл с текстом (иначе — stdin)") ap.add_argument("--model", default="fef2/ner_rus_bert-secret_detection", help="id на Hub или локальный каталог") ap.add_argument("--device", default="cpu", help="cpu | cuda | cuda:0 | mps") ap.add_argument("--rules", default=str(DEFAULT_RULES), help="путь к gitleaks.toml") ap.add_argument("--no-rules", action="store_true", help="только NER, без детерминированного слоя") ap.add_argument("--json", action="store_true", help="спаны в JSON") ap.add_argument("--spans", action="store_true", help="спаны построчно") args = ap.parse_args() if args.text is not None: text = args.text elif args.file: text = open(args.file, encoding="utf-8").read() else: text = sys.stdin.read() det = Detector(args.model, device=args.device, rules=None if args.no_rules else args.rules) if args.json: print(json.dumps([asdict(s) for s in det.spans(text)], ensure_ascii=False, indent=2)) elif args.spans: for s in det.spans(text): score = " ----" if s.score is None else f"{s.score:.4f}" print(f"{s.start:>7} {s.end:>7} {s.label:<16} {s.source:<5} {score} {s.text!r}") else: print(det.mask(text)) return 0 if __name__ == "__main__": raise SystemExit(main())