FinSight / backend /relation_extractor.py
Sanjam19's picture
Deploy FinSight demo (single-container Docker Space)
d4f8959 verified
Raw
History Blame Contribute Delete
9.8 kB
# backend/relation_extractor.py
import json
import os
import re
import requests
from dotenv import load_dotenv
load_dotenv()
# same env vars as backend/llm.py so one setting controls every local call
# (the old hardcoded mistral:7b default silently 404'd on machines that
# never pulled that exact model — the fallback looked enabled but never ran)
_OLLAMA_BASE = os.getenv("OLLAMA_URL", "http://localhost:11434")
OLLAMA_URL = f"{_OLLAMA_BASE}/api/generate"
OLLAMA_MODEL = os.getenv("OLLAMA_MODEL", "qwen2.5:1.5b")
OLLAMA_TIMEOUT = 60 # seconds per call — CPU inference, not GPU
# ===== Pattern-based extraction (runs first, free, instant) =====
# Each pattern: (regex, rel_type). {A} and {B} are org mention placeholders
# already substituted before matching — see extract_patterns().
RELATION_PATTERNS = [
(r"\bsubsidiary of\b", "SUBSIDIARY_OF"),
(r"\bwholly[\s-]owned subsidiary of\b", "SUBSIDIARY_OF"),
(r"\bparent company\b", "SUBSIDIARY_OF"),
(r"\bcompetes? with\b", "COMPETITOR_OF"),
(r"\bcompetitor[s]? (?:of|to|include[s]?)\b", "COMPETITOR_OF"),
(r"\brivals?\b", "COMPETITOR_OF"),
(r"\bsupplies? (?:to|for)\b", "SUPPLIER_TO"),
(r"\bsupplier (?:of|to|for)\b", "SUPPLIER_TO"),
(r"\bvendor (?:of|to|for)\b", "SUPPLIER_TO"),
(r"\bpartnered? with\b", "PARTNERED_WITH"),
(r"\bpartnership with\b", "PARTNERED_WITH"),
(r"\bjoint venture with\b", "PARTNERED_WITH"),
(r"\bcollaborat\w* with\b", "PARTNERED_WITH"),
(r"\bacquired\b", "ACQUIRED"),
(r"\bacquisition of\b", "ACQUIRED"),
]
def _find_mentions(text: str, org_texts: list) -> list:
"""Find character spans of each org mention in text. Returns list of
(start, end, org_text), sorted by position, with overlapping spans
collapsed to the LONGEST match (so 'Apple' inside 'Apple Inc.' doesn't
also register as a separate standalone mention)."""
raw_spans = []
for org in org_texts:
for m in re.finditer(re.escape(org), text, re.IGNORECASE):
raw_spans.append((m.start(), m.end(), org))
raw_spans.sort(key=lambda s: (s[0], -(s[1] - s[0]))) # longest first per start
spans = []
for start, end, org in raw_spans:
overlaps = any(
not (end <= s_start or start >= s_end)
for s_start, s_end, _ in spans
)
if not overlaps:
spans.append((start, end, org))
spans.sort(key=lambda s: s[0])
return spans
def extract_patterns(text: str, org_texts: list) -> list:
"""Pattern-based pass. Returns list of (org_a, org_b, rel_type).
Only considers org pairs that co-occur within the SAME SENTENCE —
char-distance alone isn't enough since two unrelated mentions in
adjacent sentences can still fall inside a tight char window."""
spans = _find_mentions(text, org_texts)
# map each span to the sentence index it falls in
sentence_bounds = []
pos = 0
for sent in re.split(r"(?<=[.!?])\s+", text):
sentence_bounds.append((pos, pos + len(sent)))
pos += len(sent) + 1 # approximate, accounts for the split separator
def sentence_index(char_pos):
for idx, (s_start, s_end) in enumerate(sentence_bounds):
if s_start <= char_pos < s_end:
return idx
return len(sentence_bounds) # past the end, treat as unique bucket
found = {} # pair_key -> (org_a, org_b, rel_type), first match wins
for i in range(len(spans)):
for j in range(i + 1, len(spans)):
start_a, end_a, org_a = spans[i]
start_b, end_b, org_b = spans[j]
if org_a.lower() == org_b.lower():
continue
if sentence_index(start_a) != sentence_index(start_b):
continue # different sentences, don't connect them
gap_start, gap_end = end_a, start_b
if gap_end < gap_start:
continue
between = text[gap_start:gap_end].lower()
matched_rel = None
for pattern, rel_type in RELATION_PATTERNS:
if re.search(pattern, between):
matched_rel = rel_type
break
if matched_rel:
pair_key = frozenset([org_a.lower(), org_b.lower()])
if pair_key not in found:
found[pair_key] = (org_a, org_b, matched_rel)
return list(found.values())
# ===== LLM fallback (local Ollama, only for unmatched co-occurring pairs) =====
LLM_PROMPT_TEMPLATE = """You are extracting relationships between two companies mentioned in a financial document excerpt.
Excerpt:
\"\"\"{chunk_text}\"\"\"
Company A: {org_a}
Company B: {org_b}
Based ONLY on the excerpt above, what is the relationship between Company A and Company B?
Choose exactly one label from this list: SUBSIDIARY_OF, COMPETITOR_OF, SUPPLIER_TO, PARTNERED_WITH, ACQUIRED, BOARD_OVERLAP_WITH, NONE
Respond with ONLY a JSON object, nothing else, in this exact format:
{{"relation": "LABEL", "confidence": "high|low"}}
If the excerpt does not clearly support a relationship, respond with {{"relation": "NONE", "confidence": "low"}}.
"""
VALID_LLM_RELATIONS = {
"SUBSIDIARY_OF", "COMPETITOR_OF", "SUPPLIER_TO",
"PARTNERED_WITH", "ACQUIRED", "BOARD_OVERLAP_WITH"
}
def _ollama_available() -> bool:
try:
r = requests.get(f"{_OLLAMA_BASE}/api/tags", timeout=2)
return r.status_code == 200
except requests.RequestException:
return False
def extract_llm(chunk_text: str, org_a: str, org_b: str) -> dict | None:
"""Single LLM call for one ambiguous pair. Returns
{"relation": ..., "confidence": ...} or None on any failure
(Ollama down, bad JSON, invalid label, timeout)."""
prompt = LLM_PROMPT_TEMPLATE.format(
chunk_text=chunk_text[:1000], # keep prompt short, local model
org_a=org_a,
org_b=org_b
)
payload = {
"model": OLLAMA_MODEL,
"prompt": prompt,
"stream": False,
"options": {
"temperature": 0,
"num_ctx": int(os.getenv("OLLAMA_NUM_CTX", "4096"))
}
}
# suppress chain-of-thought for thinking models (mirrors backend/llm.py)
if OLLAMA_MODEL.split(":")[0] in ("qwen3", "deepseek-r1"):
payload["think"] = False
try:
resp = requests.post(
OLLAMA_URL,
json=payload,
timeout=OLLAMA_TIMEOUT
)
resp.raise_for_status()
raw = resp.json().get("response", "").strip()
# strip thinking blocks and markdown fences if the model adds them anyway
raw = re.sub(r"<think>.*?</think>", "", raw, flags=re.DOTALL).strip()
raw = re.sub(r"^```(?:json)?|```$", "", raw, flags=re.MULTILINE).strip()
parsed = json.loads(raw)
relation = parsed.get("relation", "NONE")
if relation not in VALID_LLM_RELATIONS:
return None
return {
"relation": relation,
"confidence": parsed.get("confidence", "low")
}
except (requests.RequestException, json.JSONDecodeError, ValueError, KeyError):
return None
# ===== Main entry point =====
def extract_relations(
chunk_text: str,
org_texts: list,
use_llm_fallback: bool = True
) -> list:
"""
Extract typed relationships between co-occurring ORG entities in a chunk.
Args:
chunk_text: full text of the chunk
org_texts: list of ORG entity mention strings found in this chunk
use_llm_fallback: if True, calls local Ollama for pairs the
regex patterns miss. Silently skipped if Ollama isn't running.
Returns:
list of dicts: {"org_a": str, "org_b": str, "relation": str, "source": "pattern"|"llm"}
"""
if len(org_texts) < 2:
return []
results = []
pattern_hits = extract_patterns(chunk_text, org_texts)
matched_pairs = set()
for org_a, org_b, rel_type in pattern_hits:
results.append({
"org_a": org_a, "org_b": org_b,
"relation": rel_type, "source": "pattern"
})
matched_pairs.add(frozenset([org_a.lower(), org_b.lower()]))
if not use_llm_fallback:
return results
spans = _find_mentions(chunk_text, org_texts)
unique_orgs = list({s[2] for s in spans})
if len(unique_orgs) < 2:
return results
if not _ollama_available():
print("relation_extractor: Ollama not reachable, skipping LLM fallback")
return results
for i in range(len(unique_orgs)):
for j in range(i + 1, len(unique_orgs)):
org_a, org_b = unique_orgs[i], unique_orgs[j]
if org_a.lower() == org_b.lower():
continue
pair_key = frozenset([org_a.lower(), org_b.lower()])
if pair_key in matched_pairs:
continue # pattern already found a relation for this pair
llm_result = extract_llm(chunk_text, org_a, org_b)
if llm_result and llm_result["relation"] != "NONE":
results.append({
"org_a": org_a, "org_b": org_b,
"relation": llm_result["relation"],
"source": "llm",
"confidence": llm_result["confidence"]
})
return results
if __name__ == "__main__":
test_chunk = (
"Foxconn is a major supplier to Apple Inc. for iPhone assembly. "
"In contrast, Apple competes with Samsung in the smartphone market. "
"Beats Electronics, a subsidiary of Apple, also contributed to revenue."
)
test_orgs = ["Foxconn", "Apple Inc.", "Samsung", "Beats Electronics", "Apple"]
rels = extract_relations(test_chunk, test_orgs, use_llm_fallback=True)
for r in rels:
print(r)