Spaces:
Runtime error
Runtime error
File size: 44,360 Bytes
0992cee 97ff337 74fe989 0992cee 0f49634 0992cee 97ff337 0992cee 97c3931 0992cee 0f49634 0992cee 97ff337 0992cee 97ff337 0992cee 74fe989 0992cee 97ff337 0992cee 97ff337 97c3931 74fe989 97c3931 97ff337 97c3931 97ff337 97c3931 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97c3931 0992cee 97c3931 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 97ff337 0992cee 0f49634 0992cee 0f49634 0992cee 97ff337 0f49634 97ff337 0f49634 97ff337 0f49634 97ff337 0f49634 97ff337 0992cee 97ff337 0992cee 97c3931 97ff337 74fe989 0992cee | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 | """Optional knowledge-graph + embedding search over EXTRACTED text (Gemini).
After a PDF is extracted, this module turns its per-page text into a small
**knowledge graph** β typed entities (nodes) joined by ``(subject, predicate,
object)`` triples (edges) β embeds every node into a vector space, and answers a
query by combining **semantic vector search** (conceptually-related entities)
with **explicit graph traversal** (following the facts). The answer is
*explainable*: every result carries the supporting triples and the page each
came from.
DESIGN / IDENTITY
- This is an OPT-IN, online feature. Like ``online_ocr.py`` it sends text to
Google's Gemini API, so it is OFF by default and only runs when the caller
passes a Gemini API key. The page text leaves this machine.
- ZERO new pip dependencies: the REST calls use the Python standard library
only (``urllib.request`` + ``json``); the graph is a plain in-memory
adjacency structure; ``numpy`` (already a core dep) holds the node vectors.
This keeps the feature working on the lean hosted build too (no torch /
sentence-transformers / networkx required).
- The security-critical key sanitization and the HTTP error mapping are REUSED
from ``online_ocr`` so the "never echo the key, only raise a single-line
RuntimeError" contract holds here as well.
Importing this module performs NO network call and has no import-time side
effects.
Public API
----------
``extract_triples(text, *, api_key, model=None, timeout=...) -> dict``
Pull ``{"entities": [...], "triples": [...]}`` out of one page of text.
``embed_texts(texts, *, api_key, model=None, task_type=..., timeout=...) -> np.ndarray``
Embed a list of strings to an ``(n, dim)`` float32 matrix.
``build_graph(pages, *, api_key, ...) -> KnowledgeGraph``
Build a per-document graph from ``[(page_number, text), ...]``.
``KnowledgeGraph``
Holds the graph + node vectors; ``.search(query, ...)`` does hybrid
retrieval and returns ranked, explainable results.
"""
import hashlib
import http.client
import json
import re
import time
import urllib.error
import urllib.request
from collections import deque
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from typing import Optional
import numpy as np
# Reuse the verified, security-tested helpers from the OCR client so the key is
# sanitized identically and HTTP errors never leak it.
from pipeline.online_ocr import (
_API_BASE,
_API_KEY_HEADER,
_clean_key,
_friendly_http_error,
is_configured,
)
# ---------------------------------------------------------------------------
# Gemini REST endpoints (same base/host as online_ocr; different methods).
# ---------------------------------------------------------------------------
_GENERATE_URL = _API_BASE + "/{model}:generateContent"
_BATCH_EMBED_URL = _API_BASE + "/{model}:batchEmbedContents"
_EMBED_URL = _API_BASE + "/{model}:embedContent" # single-item fallback
# Triple extraction must default to a FREE-TIER model (Flash). Pro/preview
# models return HTTP 429 limit:0 on the free tier β see online_ocr notes.
DEFAULT_TRIPLE_MODEL = "gemini-2.5-flash"
# Embeddings: the wire format wants the "models/" prefix in the per-request
# "model" field. The embedding model id has CHANGED over time β Google retired
# `text-embedding-004` on v1beta (it now 404s) in favour of `gemini-embedding-001`.
# So we DON'T trust a single hard-coded id: build_graph DISCOVERS the model the
# key can actually use via ModelService.ListModels (`_pick_embedding_model`),
# trying this preference order and falling back to the default only if discovery
# fails. Order newestβoldest so a key that still has an old model also works.
DEFAULT_EMBED_MODEL = "gemini-embedding-001"
EMBED_MODEL_PREFERENCE = (
"gemini-embedding-001", # current GA model (default ~3072-dim; we L2-normalize)
"text-embedding-004", # legacy 768-dim (retired on v1beta for many keys)
"embedding-001", # older fallback
)
# Gemini caps batchEmbedContents at 100 requests per call; chunk to stay under.
_EMBED_BATCH = 100
# Defensive cap on per-page text sent for triple extraction. A single PDF page
# is tiny next to Flash's context window, but a pathological page (e.g. a giant
# embedded text dump) shouldn't balloon the request.
_MAX_PAGE_CHARS = 100_000
# Transient-failure backoff. Free-tier embedding/generate limits are real, so a
# 429/503 is retried a few times with exponential backoff before giving up.
_RETRY_STATUSES = frozenset({429, 503})
_RETRY_ATTEMPTS = 4
_RETRY_BASE_DELAY = 1.5 # seconds; doubled each attempt
# Build scalability: per-page triple extraction is one (I/O-bound) Gemini call
# each, so they're fanned across this SHARED, bounded pool instead of running
# strictly sequentially β a multi-page build is ~Nx faster on a key with real
# throughput. The pool is module-level on purpose: being shared, it also caps the
# TOTAL number of concurrent Gemini calls across simultaneous builds/users, so one
# big document can't blow the rate limit for everyone. Overshoot of the provider's
# rate limit is absorbed by the 429/503 backoff in _post_json. Graph mutation stays
# single-threaded (results are collected, then applied in deterministic page order).
_BUILD_CONCURRENCY = 5
_EXTRACT_POOL = ThreadPoolExecutor(
max_workers=_BUILD_CONCURRENCY, thread_name_prefix="kg-extract"
)
# Instruction for triple extraction. We ask for typed entities AND triples and
# force a JSON response via responseSchema for deterministic parsing.
_TRIPLE_PROMPT = (
"You are building a knowledge graph from a document. From the TEXT below, "
"extract the key entities and the factual relationships between them.\n"
"- entities: the important named things (people, organizations, places, "
"dates, IDs, amounts, concepts). Give each a short canonical name and a "
"TYPE (e.g. PERSON, ORG, PLACE, DATE, ID, AMOUNT, CONCEPT, OTHER).\n"
"- triples: factual (subject, predicate, object) statements stated or "
"clearly implied by the text. Use concise predicates (e.g. 'works at', "
"'located in', 'has id', 'dated'). Subject and object should be entity "
"names. Do NOT invent facts that are not supported by the text.\n"
"Return only what the text supports; if the text is empty or has no "
"extractable facts, return empty lists.\n\nTEXT:\n"
)
# Response schema for generateContent β keeps output a strict JSON object.
_TRIPLE_SCHEMA = {
"type": "OBJECT",
"properties": {
"entities": {
"type": "ARRAY",
"items": {
"type": "OBJECT",
"properties": {
"name": {"type": "STRING"},
"type": {"type": "STRING"},
},
"required": ["name"],
},
},
"triples": {
"type": "ARRAY",
"items": {
"type": "OBJECT",
"properties": {
"subject": {"type": "STRING"},
"predicate": {"type": "STRING"},
"object": {"type": "STRING"},
},
"required": ["subject", "predicate", "object"],
},
},
},
"required": ["triples"],
}
# ---------------------------------------------------------------------------
# Low-level HTTP (stdlib) β mirrors online_ocr's request/error handling.
# ---------------------------------------------------------------------------
class GeminiHTTPError(RuntimeError):
"""A friendly, KEY-FREE Gemini HTTP error that also carries the status code.
Subclasses RuntimeError so every existing ``except RuntimeError`` still
catches it; the ``.code`` lets callers branch (e.g. fall back from
batchEmbedContents to embedContent on a 404 method-not-found).
"""
def __init__(self, message, code=None):
super().__init__(message)
self.code = code
def _post_json(url: str, body: dict, api_key: str, *, timeout: float) -> dict:
"""POST a JSON body to Gemini and return the parsed JSON response.
``api_key`` must already be sanitized via ``_clean_key``. Retries 429/503
with exponential backoff, then maps any HTTP/network/non-JSON failure to a
single-line RuntimeError (never echoing the key β ``_friendly_http_error``
reads only the API's own error body).
"""
data = json.dumps(body).encode("utf-8")
request = urllib.request.Request(
url,
data=data,
method="POST",
headers={
"Content-Type": "application/json",
_API_KEY_HEADER: api_key,
},
)
last_http_err: Optional[urllib.error.HTTPError] = None
for attempt in range(_RETRY_ATTEMPTS):
try:
with urllib.request.urlopen(request, timeout=timeout) as response:
raw = response.read()
break
except urllib.error.HTTPError as err:
last_http_err = err
if err.code in _RETRY_STATUSES and attempt < _RETRY_ATTEMPTS - 1:
time.sleep(_RETRY_BASE_DELAY * (2 ** attempt))
continue
raise GeminiHTTPError(str(_friendly_http_error(err)), code=err.code) from None
except urllib.error.URLError as err:
raise RuntimeError(
"Could not reach the Gemini API ({}). Check your internet "
"connection.".format(err.reason)
) from None
except TimeoutError:
raise RuntimeError(
"The Gemini request timed out after {}s.".format(timeout)
) from None
except (OSError, http.client.HTTPException) as err:
# The connection dropped mid-body: urlopen() returned headers but
# response.read() failed (e.g. ConnectionResetError β an OSError not
# wrapped in URLError β or http.client.IncompleteRead, which isn't
# even an OSError). This clause MUST sit after URLError/TimeoutError
# (both OSError subclasses). Mapping it to a single-line RuntimeError
# keeps the "_post_json only ever raises RuntimeError" contract that
# build_graph's per-page isolation and non-fatal embedding rely on.
raise RuntimeError(
"The Gemini connection dropped while reading the response ({}); "
"retry shortly.".format(err)
) from None
else: # pragma: no cover - loop always breaks or raises
raise GeminiHTTPError(
str(_friendly_http_error(last_http_err)),
code=getattr(last_http_err, "code", None),
)
try:
payload = json.loads(raw.decode("utf-8", "replace"))
except (ValueError, AttributeError) as err:
raise RuntimeError(
"Gemini returned a response that was not valid JSON ({}). The "
"service may be having problems; retry shortly.".format(err)
) from None
if not isinstance(payload, dict):
raise RuntimeError(
"Gemini returned an unexpected response (not a JSON object); "
"retry shortly."
)
return payload
def _resolve_model(model, default: str) -> str:
"""Return a clean bare model id (no ``models/`` prefix), defaulting when unset."""
name = (str(model).strip() if model else "") or default
if name.startswith("models/"):
name = name[len("models/"):]
return name or default
# Process-level cache of the discovered embedding model, keyed by a HASH of the
# API key (never the key itself). Different keys/projects can have access to
# different models β on the multi-user demo Space the first visitor's key must
# NOT pin the model for everyone. Only a SUCCESSFUL discovery is cached, so a
# transient failure that fell back to the default is retried next time.
_EMBED_MODEL_RESOLVED: dict = {}
def _list_embedding_models(api_key, *, timeout=30) -> list:
"""Return bare ids of models the key can use for embedding.
Calls ``GET /v1beta/models`` and keeps those whose
``supportedGenerationMethods`` include ``batchEmbedContents`` (what we call)
or ``embedContent``. Used to pick a model that actually exists, instead of
hard-coding an id that Google may retire (the cause of the text-embedding-004
404). Mirrors ``online_ocr.list_models`` but filters for embedding support.
"""
api_key = _clean_key(api_key)
request = urllib.request.Request(
_API_BASE, method="GET", headers={_API_KEY_HEADER: api_key}
)
try:
with urllib.request.urlopen(request, timeout=timeout) as response:
raw = response.read()
except urllib.error.HTTPError as err:
raise _friendly_http_error(err) from None
except urllib.error.URLError as err:
raise RuntimeError(
"Could not reach the Gemini API ({}). Check your internet "
"connection.".format(err.reason)
) from None
except TimeoutError:
raise RuntimeError(
"The Gemini request timed out after {}s.".format(timeout)
) from None
except (OSError, http.client.HTTPException) as err:
# Connection dropped while reading the ListModels body (see _post_json).
# Keep it a single-line RuntimeError so _pick_embedding_model's
# `except RuntimeError` catches it and falls back to the default model.
raise RuntimeError(
"The Gemini connection dropped while reading the response ({}); "
"retry shortly.".format(err)
) from None
try:
payload = json.loads(raw.decode("utf-8", "replace"))
except (ValueError, AttributeError):
return []
out = []
for entry in (payload.get("models") or []):
if not isinstance(entry, dict):
continue
methods = entry.get("supportedGenerationMethods") or []
if "batchEmbedContents" in methods or "embedContent" in methods:
name = str(entry.get("name") or "")
if name.startswith("models/"):
name = name[len("models/"):]
if name:
out.append(name)
return out
def _pick_embedding_model(api_key, requested=None, *, timeout=30) -> str:
"""Choose an embedding model id the key can actually use.
An explicit ``requested`` model wins. Otherwise discover the available
embedding models and pick by :data:`EMBED_MODEL_PREFERENCE`; on a transient
discovery failure fall back to :data:`DEFAULT_EMBED_MODEL` (without caching,
so it's retried). A successful discovery is cached for the process.
"""
if requested:
return _resolve_model(requested, DEFAULT_EMBED_MODEL)
key_hash = hashlib.sha256(_clean_key(api_key).encode("utf-8")).hexdigest()
cached = _EMBED_MODEL_RESOLVED.get(key_hash)
if cached:
return cached
try:
available = set(_list_embedding_models(api_key, timeout=timeout))
except RuntimeError:
available = set()
chosen = next((m for m in EMBED_MODEL_PREFERENCE if m in available), None)
if not chosen and available:
# An unfamiliar but embedding-capable model β prefer one that looks like
# an embedding model, else just take the first deterministically.
chosen = next(
(m for m in sorted(available) if "embedding" in m.lower()),
sorted(available)[0],
)
if chosen:
_EMBED_MODEL_RESOLVED[key_hash] = chosen # cache only a real discovery
return chosen
return DEFAULT_EMBED_MODEL
# ---------------------------------------------------------------------------
# Response parsers (pure functions β unit-tested with canned dicts, no network).
# ---------------------------------------------------------------------------
def _response_text(payload: dict) -> str:
"""Concatenate the text parts of a generateContent response.
Handles a prompt-level block or a non-STOP finishReason by raising an
actionable RuntimeError, and tolerates an explicit ``{"text": null}`` part.
"""
feedback = payload.get("promptFeedback") or {}
block_reason = feedback.get("blockReason")
candidates = payload.get("candidates") or []
if not candidates:
if block_reason:
raise RuntimeError(
"Gemini blocked the request (blockReason={}). The text tripped "
"a safety filter; knowledge-graph extraction is unavailable for "
"this page.".format(block_reason)
)
raise RuntimeError(
"Gemini returned no candidates for knowledge-graph extraction."
)
candidate = candidates[0] or {}
content = candidate.get("content") or {}
parts = content.get("parts") or []
return "".join(
str(part.get("text") or "") for part in parts if isinstance(part, dict)
)
def _parse_triple_payload(payload: dict) -> dict:
"""Parse a generateContent response into ``{"entities": [...], "triples": [...]}``.
The model is asked to return a JSON object (responseMimeType=application/json),
so the candidate text is itself a JSON string. We parse it and coerce every
field to clean strings, dropping malformed/empty rows. Never raises on a
merely-empty or slightly-malformed result β returns empty lists instead, so
one odd page can't fail the whole build.
"""
text = _response_text(payload).strip()
if not text:
return {"entities": [], "triples": []}
# Strip a ```json ... ``` fence if the model wrapped its JSON in one β but
# only the fence markers, so backticks legitimately inside the content
# (e.g. a value with a code span) are preserved.
if text.startswith("```"):
text = text[3:]
if text[:4].lower() == "json":
text = text[4:]
if text.endswith("```"):
text = text[:-3]
text = text.strip()
try:
obj = json.loads(text)
except (ValueError, TypeError):
return {"entities": [], "triples": []}
if not isinstance(obj, dict):
# Some models return a bare array of triples.
obj = {"triples": obj} if isinstance(obj, list) else {}
entities = []
for e in obj.get("entities") or []:
if not isinstance(e, dict):
continue
name = str(e.get("name") or "").strip()
if not name:
continue
etype = str(e.get("type") or "").strip().upper() or "OTHER"
entities.append({"name": name, "type": etype})
triples = []
for t in obj.get("triples") or []:
if not isinstance(t, dict):
continue
subj = str(t.get("subject") or "").strip()
pred = str(t.get("predicate") or "").strip()
obj_ = str(t.get("object") or "").strip()
if subj and pred and obj_:
triples.append({"subject": subj, "predicate": pred, "object": obj_})
return {"entities": entities, "triples": triples}
def _parse_embed_payload(payload: dict, expected: int) -> list:
"""Parse a batchEmbedContents response into a list of float vectors.
Response shape: ``{"embeddings": [{"values": [...]}, ...]}`` in request
order. Raises if the count doesn't match what we sent (a silent mismatch
would misalign vectors with nodes).
"""
embeddings = payload.get("embeddings")
if not isinstance(embeddings, list):
raise RuntimeError(
"Gemini embedding response had no 'embeddings' list; cannot build "
"the vector index."
)
if len(embeddings) != expected:
raise RuntimeError(
"Gemini returned {} embeddings for {} inputs (count mismatch).".format(
len(embeddings), expected
)
)
out = []
for item in embeddings:
values = (item or {}).get("values") if isinstance(item, dict) else None
if not isinstance(values, list) or not values:
raise RuntimeError("Gemini returned an empty embedding vector.")
out.append([float(v) for v in values])
return out
def _parse_single_embed_payload(payload: dict) -> list:
"""Parse a single-item embedContent response: ``{"embedding": {"values": [...]}}``."""
emb = payload.get("embedding") if isinstance(payload, dict) else None
values = emb.get("values") if isinstance(emb, dict) else None
if not isinstance(values, list) or not values:
raise RuntimeError("Gemini returned an empty embedding vector.")
return [float(v) for v in values]
def _embed_item(wire_model: str, text: str, task_type: str) -> dict:
"""One embedding request item (same shape for batch list and single body)."""
return {
"model": wire_model,
"content": {"parts": [{"text": text}]},
"taskType": task_type,
}
# ---------------------------------------------------------------------------
# Public extraction / embedding calls.
# ---------------------------------------------------------------------------
def extract_triples(text: str, *, api_key, model=None, timeout: float = 120) -> dict:
"""Extract typed entities + ``(subject, predicate, object)`` triples from text.
Returns ``{"entities": [{"name","type"}], "triples":
[{"subject","predicate","object"}]}``. Provenance (the source page) is NOT
returned by the model β the caller attaches it.
"""
if not is_configured(api_key):
raise RuntimeError(
"Knowledge-graph extraction needs a Gemini API key but none was "
"provided. Get a free key from https://aistudio.google.com/."
)
clean = str(text or "").strip()
if not clean:
return {"entities": [], "triples": []}
if len(clean) > _MAX_PAGE_CHARS:
clean = clean[:_MAX_PAGE_CHARS]
api_key = _clean_key(api_key)
model_id = _resolve_model(model, DEFAULT_TRIPLE_MODEL)
body = {
"contents": [{"parts": [{"text": _TRIPLE_PROMPT + clean}]}],
"generationConfig": {
"temperature": 0,
"responseMimeType": "application/json",
"responseSchema": _TRIPLE_SCHEMA,
},
}
payload = _post_json(
_GENERATE_URL.format(model=model_id), body, api_key, timeout=timeout
)
return _parse_triple_payload(payload)
def embed_texts(
texts,
*,
api_key,
model=None,
task_type: str = "RETRIEVAL_DOCUMENT",
timeout: float = 120,
) -> np.ndarray:
"""Embed a list of strings into an ``(n, dim)`` float32 matrix via Gemini.
``task_type`` should be ``RETRIEVAL_DOCUMENT`` for the graph's nodes and
``RETRIEVAL_QUERY`` for the search query β the embedding model uses it to
tune the vectors for retrieval, a real accuracy win. Batches at 100 inputs
per request (Gemini's cap) and concatenates in order. Pass an explicit
``model`` (typically resolved by :func:`_pick_embedding_model`); the default
is only a last resort since Google retires embedding-model ids over time.
"""
if not is_configured(api_key):
raise RuntimeError(
"Knowledge-graph search needs a Gemini API key but none was provided."
)
items = [str(t or "") for t in texts]
if not items:
return np.zeros((0, 0), dtype=np.float32)
api_key = _clean_key(api_key)
model_id = _resolve_model(model, DEFAULT_EMBED_MODEL)
wire_model = "models/" + model_id
batch_url = _BATCH_EMBED_URL.format(model=model_id)
single_url = _EMBED_URL.format(model=model_id)
vectors: list = []
use_batch = True
for start in range(0, len(items), _EMBED_BATCH):
chunk = items[start:start + _EMBED_BATCH]
if use_batch:
body = {"requests": [_embed_item(wire_model, t, task_type) for t in chunk]}
try:
payload = _post_json(batch_url, body, api_key, timeout=timeout)
vectors.extend(_parse_embed_payload(payload, len(chunk)))
continue
except GeminiHTTPError as exc:
# 404 = this model/version doesn't expose batchEmbedContents.
# Degrade to per-item embedContent instead of failing the build.
if exc.code != 404:
raise
use_batch = False
for t in chunk:
payload = _post_json(single_url, _embed_item(wire_model, t, task_type),
api_key, timeout=timeout)
vectors.append(_parse_single_embed_payload(payload))
return np.asarray(vectors, dtype=np.float32)
# ---------------------------------------------------------------------------
# Graph data structures + builder.
# ---------------------------------------------------------------------------
def _norm_key(name: str) -> str:
"""Canonical key for entity de-duplication (case/space-insensitive)."""
return re.sub(r"\s+", " ", str(name or "").strip()).casefold()
_WORD_RE = re.compile(r"[^\W\d_]+|\d+", re.UNICODE)
def _tokens(text: str) -> set:
return {t for t in _WORD_RE.findall(str(text or "").casefold()) if len(t) > 1}
@dataclass
class KnowledgeGraph:
"""A per-document knowledge graph + node embeddings, held in RAM.
Nodes are entities; edges are ``(subject, predicate, object)`` triples, each
carrying the page it came from (provenance). Node vectors live in a parallel
float32 matrix (L2-normalized for cosine). Everything is plain Python +
numpy β no networkx, no external store β so it evicts with its ``Job`` and
serializes only via the explicit ``to_*`` methods (never leaked by accident).
"""
# Parallel arrays indexed by node id (0..n-1).
nodes: list = field(default_factory=list) # [{name, type, pages, degree}]
triples: list = field(default_factory=list) # [{subject_id, predicate, object_id, subject, object, page}]
adjacency: dict = field(default_factory=dict) # node_id -> [(neighbor_id, triple_index)]
_index: dict = field(default_factory=dict, repr=False) # norm_key -> node_id
vectors: Optional[np.ndarray] = field(default=None, repr=False) # (n, dim) L2-normalized
embed_model: Optional[str] = None
# Build telemetry (surfaced so partial builds / degraded search aren't silent).
pages_built: int = 0 # pages whose triple extraction succeeded
pages_failed: int = 0 # pages skipped due to a per-page Gemini failure
embed_error: Optional[str] = None # set if embedding failed -> lexical-only
# --- construction helpers ------------------------------------------------
def _add_node(self, name: str, etype: str = "OTHER", page: Optional[int] = None) -> int:
key = _norm_key(name)
if not key:
return -1
nid = self._index.get(key)
if nid is None:
nid = len(self.nodes)
self._index[key] = nid
self.nodes.append(
{"name": str(name).strip(), "type": (etype or "OTHER"), "pages": [], "degree": 0}
)
self.adjacency[nid] = []
node = self.nodes[nid]
# Prefer a known type over OTHER if a later mention supplies one.
if (not node["type"] or node["type"] == "OTHER") and etype and etype != "OTHER":
node["type"] = etype
if page is not None and page not in node["pages"]:
node["pages"].append(page)
return nid
def _add_triple(self, subject: str, predicate: str, obj: str, page: Optional[int]) -> None:
s_id = self._add_node(subject, page=page)
o_id = self._add_node(obj, page=page)
if s_id < 0 or o_id < 0 or s_id == o_id:
return
tindex = len(self.triples)
self.triples.append(
{
"subject_id": s_id,
"object_id": o_id,
"predicate": str(predicate).strip(),
"subject": self.nodes[s_id]["name"],
"object": self.nodes[o_id]["name"],
"page": page,
}
)
# Undirected adjacency for traversal; direction is kept in the triple.
self.adjacency[s_id].append((o_id, tindex))
self.adjacency[o_id].append((s_id, tindex))
self.nodes[s_id]["degree"] += 1
self.nodes[o_id]["degree"] += 1
def _node_embed_text(self, nid: int) -> str:
"""Text used to embed a node: its name/type + a little fact context."""
node = self.nodes[nid]
bits = ["{} ({})".format(node["name"], node["type"])]
# Up to a few connected facts give the embedding semantic context.
ctx = [
"{} {} {}".format(t["subject"], t["predicate"], t["object"])
for _, tindex in self.adjacency[nid][:3]
for t in (self.triples[tindex],)
]
if ctx:
bits.append(". ".join(ctx))
return ". ".join(bits)
def set_vectors(self, vectors: Optional[np.ndarray], model: Optional[str] = None) -> None:
"""Attach (and L2-normalize) the node-vector matrix."""
if vectors is None or getattr(vectors, "size", 0) == 0:
self.vectors = None
return
v = np.asarray(vectors, dtype=np.float32)
norms = np.linalg.norm(v, axis=1, keepdims=True)
norms[norms == 0] = 1.0
self.vectors = v / norms
self.embed_model = model
# --- public properties ---------------------------------------------------
@property
def num_nodes(self) -> int:
return len(self.nodes)
@property
def num_edges(self) -> int:
return len(self.triples)
@property
def has_vectors(self) -> bool:
return self.vectors is not None and self.vectors.shape[0] == len(self.nodes)
# --- retrieval -----------------------------------------------------------
def _semantic_scores(self, query: str, query_vec: Optional[np.ndarray]):
"""Per-node relevance to the query in [0,1].
Returns ``(scores, used_vectors)``. Uses cosine similarity when vectors +
a usable query vector are available; otherwise falls back to lexical
token overlap (``used_vectors=False``) so the caller can report the mode
accurately even when it silently fell back (e.g. a dim mismatch).
"""
n = len(self.nodes)
if n == 0:
return np.zeros((0,), dtype=np.float32), False
if self.has_vectors and query_vec is not None and getattr(query_vec, "size", 0):
q = np.asarray(query_vec, dtype=np.float32).reshape(-1)
qn = np.linalg.norm(q)
if qn > 0 and q.shape[0] == self.vectors.shape[1]:
sims = self.vectors @ (q / qn) # cosine; both are L2-normalized
return ((sims + 1.0) / 2.0).astype(np.float32), True # [-1,1]->[0,1]
# Lexical fallback.
q_tokens = _tokens(query)
scores = np.zeros((n,), dtype=np.float32)
if not q_tokens:
return scores, False
for i, node in enumerate(self.nodes):
nt = _tokens(node["name"])
if nt:
scores[i] = len(q_tokens & nt) / float(len(q_tokens | nt))
return scores, False
def _multi_source_paths(self, seeds: list, hops: int) -> dict:
"""Multi-source BFS from ``seeds`` up to ``hops``.
Returns ``{node_id: (distance, source_seed_id, [triple_index,...])}``
where the triple list is the chain of edges from the nearest seed to the
node (empty for the seeds themselves).
"""
result: dict = {}
prev: dict = {}
queue: deque = deque()
for s in seeds:
if s not in result:
result[s] = (0, s, [])
queue.append(s)
while queue:
u = queue.popleft()
dist, src, _ = result[u]
if dist >= hops:
continue
for nb_id, tindex in self.adjacency.get(u, []):
if nb_id not in result:
prev[nb_id] = (u, tindex)
result[nb_id] = (dist + 1, src, [])
queue.append(nb_id)
# Reconstruct triple chains via predecessors.
for nid in list(result.keys()):
dist, src, _ = result[nid]
chain = []
cur = nid
while cur in prev:
u, tindex = prev[cur]
chain.append(tindex)
cur = u
chain.reverse()
result[nid] = (dist, src, chain)
return result
def search(
self,
query: str,
*,
api_key=None,
query_vec: Optional[np.ndarray] = None,
embed_model=None,
top_k: int = 6,
hops: int = 2,
max_results: int = 10,
timeout: float = 60,
) -> dict:
"""Hybrid retrieval: semantic seeds + graph traversal -> explainable answers.
1. Score every node's semantic relevance to ``query`` (cosine over
embeddings, else lexical).
2. Take the ``top_k`` strongest as SEED nodes.
3. Traverse up to ``hops`` from the seeds to gather connected entities +
the linking triples.
4. Rank candidates by ``semantic + graph-proximity`` and return each
answer with its supporting path (chain of triples) and page refs.
If ``query_vec`` is given it is used directly; otherwise, when an
``api_key`` is supplied and the graph has vectors, the query is embedded
via Gemini (RETRIEVAL_QUERY). With neither, search degrades to lexical.
"""
query = str(query or "").strip()
out = {
"query": query,
"answers": [],
"seeds": [],
"triples": [],
"mode": "lexical",
}
if not query or not self.nodes:
return out
# Obtain a query vector if we can (don't fail search if embedding fails).
if query_vec is None and self.has_vectors and is_configured(api_key):
try:
mat = embed_texts(
[query],
api_key=api_key,
model=embed_model or self.embed_model,
task_type="RETRIEVAL_QUERY",
timeout=timeout,
)
if mat.size:
query_vec = mat[0]
except RuntimeError:
query_vec = None
sem, used_vectors = self._semantic_scores(query, query_vec)
out["mode"] = "semantic" if used_vectors else "lexical"
if sem.size == 0 or float(sem.max()) <= 0:
return out
# Seeds = strongest semantic matches.
order = np.argsort(-sem)
seeds = [int(i) for i in order[:max(1, top_k)] if sem[int(i)] > 0]
out["seeds"] = [self.nodes[i]["name"] for i in seeds]
reach = self._multi_source_paths(seeds, hops=max(0, hops))
# Rank reachable candidates.
W_SEM, W_GRAPH, DECAY = 0.6, 0.4, 0.55
scored = []
for nid, (dist, src, chain) in reach.items():
sim = float(sem[nid])
graph_term = float(sem[src]) * (DECAY ** dist)
score = W_SEM * sim + W_GRAPH * graph_term
scored.append((score, dist, nid, src, chain))
scored.sort(key=lambda x: (-x[0], x[1]))
used_triples: dict = {}
for score, dist, nid, src, chain in scored[:max(1, max_results)]:
node = self.nodes[nid]
path_triples = [self._triple_view(t) for t in chain]
path_pages = sorted({t["page"] for t in path_triples if t["page"] is not None})
# The entity's own incident facts make even a seed (empty path)
# explainable; cap a few of the highest-information ones.
fact_idx = [tindex for _, tindex in self.adjacency.get(nid, [])][:6]
facts = [self._triple_view(t) for t in fact_idx]
for t in list(chain) + fact_idx:
used_triples[t] = self._triple_view(t)
out["answers"].append(
{
"entity": node["name"],
"type": node["type"],
"score": round(float(score), 4),
"similarity": round(float(sem[nid]), 4),
"hops": int(dist),
"pages": sorted(node["pages"]),
"path": path_triples,
"path_pages": path_pages,
"facts": facts,
}
)
# Flat, de-duplicated list of every supporting triple referenced above.
out["triples"] = list(used_triples.values())
return out
def _triple_view(self, tindex: int) -> dict:
t = self.triples[tindex]
return {
"subject": t["subject"],
"predicate": t["predicate"],
"object": t["object"],
"page": t["page"],
}
# --- serialization (explicit; the graph is NEVER auto-serialized) --------
def to_viz_dict(self, max_nodes: int = 120) -> dict:
"""Compact node/edge lists for the vanilla-canvas visualization.
Caps to the highest-degree ``max_nodes`` so a huge graph stays drawable
(and the dropped count is reported, never silently truncated).
"""
order = sorted(
range(len(self.nodes)), key=lambda i: self.nodes[i]["degree"], reverse=True
)
keep = set(order[:max_nodes])
nodes = [
{
"id": i,
"name": self.nodes[i]["name"],
"type": self.nodes[i]["type"],
"pages": sorted(self.nodes[i]["pages"]),
"degree": self.nodes[i]["degree"],
}
for i in sorted(keep)
]
edges = [
{
"source": t["subject_id"],
"target": t["object_id"],
"predicate": t["predicate"],
"page": t["page"],
}
for t in self.triples
if t["subject_id"] in keep and t["object_id"] in keep
]
return {
"nodes": nodes,
"edges": edges,
"total_nodes": len(self.nodes),
"total_edges": len(self.triples),
"truncated": len(self.nodes) > len(keep),
}
def to_export_dict(self) -> dict:
"""Full graph for the .kg.json export (entities + triples + provenance)."""
return {
"entities": [
{"name": n["name"], "type": n["type"], "pages": sorted(n["pages"])}
for n in self.nodes
],
"triples": [self._triple_view(i) for i in range(len(self.triples))],
"node_count": len(self.nodes),
"triple_count": len(self.triples),
"embed_model": self.embed_model,
}
def build_graph(
pages,
*,
api_key,
triple_model=None,
embed_model=None,
max_pages: Optional[int] = None,
embed: bool = True,
timeout: float = 120,
) -> KnowledgeGraph:
"""Build a per-document :class:`KnowledgeGraph` from extracted page text.
Parameters
----------
pages : list[tuple[int, str]]
``(page_number, text)`` pairs β typically
``[(p.page, p.text) for p in job.pages if p.status == "done"]``.
api_key : str
A Gemini API key (required; this is the opt-in online feature).
triple_model, embed_model : str, optional
Model overrides; default to Flash + a discovered embedding model
(``_pick_embedding_model``, prefers gemini-embedding-001).
max_pages : int, optional
Cap the number of non-empty pages processed (cost control). None = all.
embed : bool
If True (default) also embed the nodes so semantic search works; if
False the graph supports lexical search only (cheaper).
Per-page triple extraction is fanned across a shared bounded pool
(``_EXTRACT_POOL``), so a multi-page build is much faster on a key with real
throughput; the shared pool also caps total concurrent Gemini calls. Each
page is isolated β a single failed page is skipped and counted
(``pages_failed``); only an all-pages failure raises. Embedding failure is
non-fatal (``embed_error`` set, lexical search retained).
Returns a graph (possibly empty if the document yields no facts). Raises a
single-line RuntimeError only for a missing key or a hard Gemini failure.
"""
if not is_configured(api_key):
raise RuntimeError(
"Building a knowledge graph needs a Gemini API key. Enter your free "
"key from https://aistudio.google.com/, or use the demo key."
)
usable = [(int(pn), str(txt or "")) for pn, txt in pages if str(txt or "").strip()]
if max_pages is not None and max_pages > 0:
usable = usable[:max_pages]
graph = KnowledgeGraph()
last_error: Optional[Exception] = None
# Fan the per-page Gemini calls across the shared pool (concurrency), then
# build the graph single-threaded below in PAGE ORDER so node ids stay
# deterministic and the adjacency dicts are never mutated from two threads.
def _extract_one(text):
return extract_triples(text, api_key=api_key, model=triple_model, timeout=timeout)
futures = [(page_no, _EXTRACT_POOL.submit(_extract_one, text)) for page_no, text in usable]
results: dict = {}
for page_no, fut in futures:
# Isolate each page: one rate-limited / safety-blocked page (common on the
# free tier) must NOT discard every other page (and its API spend).
try:
results[page_no] = fut.result()
except RuntimeError as exc:
results[page_no] = exc
last_error = exc
for page_no, _text in usable:
data = results.get(page_no)
if not isinstance(data, dict): # a skipped page (exception) or missing
graph.pages_failed += 1
continue
graph.pages_built += 1
# Register typed entities first so types are known, then triples.
for e in data.get("entities", []):
graph._add_node(e["name"], e.get("type", "OTHER"), page=page_no)
for t in data.get("triples", []):
graph._add_triple(t["subject"], t["predicate"], t["object"], page=page_no)
# If EVERY page failed, this is a systemic failure (bad key / quota / network),
# not "a document with no facts" β surface it instead of returning an empty graph.
if usable and graph.pages_failed == len(usable):
raise last_error or RuntimeError(
"Knowledge-graph extraction failed for every page."
)
if embed and graph.num_nodes:
# Resolve the embedding model the key can actually use (Google retires
# ids over time), so search uses the SAME model the nodes were built with.
# Embeddings are the SEMANTIC half: if they fail (quota / retired model),
# keep the graph for LEXICAL search rather than throwing away the (costly)
# triple extraction. The degradation is surfaced via has_vectors/embed_error.
try:
resolved = _pick_embedding_model(api_key, embed_model, timeout=min(timeout, 30))
node_texts = [graph._node_embed_text(i) for i in range(graph.num_nodes)]
vectors = embed_texts(
node_texts,
api_key=api_key,
model=resolved,
task_type="RETRIEVAL_DOCUMENT",
timeout=timeout,
)
graph.set_vectors(vectors, model=resolved)
except RuntimeError as exc:
graph.embed_error = str(exc)
# If the DISCOVERED model 404'd (Google retired the id mid-process),
# drop its per-key cache entry so the NEXT build re-discovers a
# working model instead of pinning the dead one for the whole process
# lifetime. Only for discovery (no explicit embed_model) and only on a
# model-not-found 404 β a transient 429/quota must NOT evict a good id.
if (
isinstance(exc, GeminiHTTPError)
and exc.code == 404
and not embed_model
):
key_hash = hashlib.sha256(
_clean_key(api_key).encode("utf-8")
).hexdigest()
_EMBED_MODEL_RESOLVED.pop(key_hash, None)
return graph
|