Skip to content
All library documents

Constructing a Supply-Chain Knowledge Graph From Company Filings

Code Machine Learning for Trading

Summary

The notebook demonstrates a pipeline for extracting supplier, customer, and competitor relationships from company annual filings and storing them as a knowledge graph. A local language model generates candidate subject–relationship–entity triples from filing text. The workflow then resolves entity names, checks the resulting data, and loads relationships into Neo4j in batches. It also describes using graph structure to inspect shared suppliers and supply-chain concentration.

The reported findings caution against treating graph size as a scaling benchmark: most extracted suppliers appear to be named by only one company, so selecting a fixed top group can obscure how little supplier overlap exists. Entity resolution relies on names and the set of filing companies, which can fail when a subsidiary is named instead of its registrant. The document says that cache provenance is checked, but edge precision has not been established. Graph relationships can support relational diagnostics and later tabular features, though the notebook does not demonstrate investment returns or predictive performance.

Key ideas

  • A language model can turn filing text into candidate supplier, customer, and competitor relationships.
  • Entity resolution uses canonical names and the roster of filing companies, but subsidiary names can still cause errors.
  • Batch graph loading separates extraction from database writes.
  • Shared-supplier counts reveal overlap that a fixed top-supplier ranking may conceal.
  • The extracted edges have not been shown to have established precision or investment value.

Tags

Full text
# 02_supply_chain_kg_construction.py


```py
# ---
# jupyter:
#   jupytext:
#     cell_metadata_filter: tags,-all
#     text_representation:
#       extension: .py
#       format_name: percent
#       format_version: '1.3'
#       jupytext_version: 1.19.3
#   kernelspec:
#     display_name: Python 3 (ipykernel)
#     language: python
#     name: python3
# ---

# %% [markdown]
# # Building a Supply Chain Knowledge Graph at Scale
#
# **Chapter 23: Knowledge Graphs for Financial AI**
#
# **Docker image**: `ml4t`
#
# > **Neo4j required**: The checked-in Qwen2.5 extraction cache is the default
# > production input, and the default path needs no GPU:
# > ```bash
# > docker compose --profile kg up -d neo4j
# > docker compose run --rm ml4t python 23_knowledge_graphs/02_supply_chain_kg_construction.py
# > ```
# > Set `RERUN_EXTRACTION=True` to rebuild the cache, which does need a CUDA GPU and
# > the `ml4t-gpu` image.
#
#
# This notebook demonstrates large-scale knowledge graph construction from SEC 10-K
# filings, extracting supply chain and competitive relationships from S&P 100 companies
# and loading them into a Neo4j graph database.
#
# **Learning Objectives**:
# - Build a relationship extraction pipeline using a local LLM (Qwen2.5-7B-Instruct)
# - Design an entity resolution step to normalize names across filings
# - Load extracted triples into Neo4j using efficient batch UNWIND queries
# - Visualize supply chain concentration risk as a network graph
#
# **Book Reference**: Chapter 23, Section 23.2 (Constructing Financial Knowledge Graphs)
#
# **Prerequisites**: Run `01_sp100_sec_download` first to download the SEC filings.
# Requires a live Neo4j instance. Cache regeneration additionally requires an RTX-class GPU.

# %%
"""Extract supply-chain relationships from SEC 10-K filings."""

from __future__ import annotations

import hashlib
import json
import logging
import os
import re
import tempfile
import time
from collections import Counter
from collections.abc import Generator
from dataclasses import dataclass
from pathlib import Path

import matplotlib.pyplot as plt
import polars as pl
import torch

from utils.paths import get_chapter_dir
from utils.reproducibility import set_global_seeds
from utils.style import COLORS, FIGSIZE, add_message_title, show_with_alt

logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %% tags=["parameters"]
# Production defaults. Papermill overrides them for testing.
# The staged SP100 10-K corpus ships with 601 filings across 101 companies, 2020-2025.
# MAX_COMPANIES caps the companies sent through the LLM when RERUN_EXTRACTION is True,
# keeping every filing of the alphabetically first N symbols. It does nothing on the
# cached path: those triples were extracted from the whole corpus.
MAX_COMPANIES = 0
LLM_BATCH_SIZE = 2  # Texts per GPU batch (Qwen2.5-7B uses ~14GB; 2 leaves headroom on 24GB)
RERUN_EXTRACTION = False  # Set True to force LLM re-extraction; False loads cached triples

# Local HuggingFace checkpoint used only for explicit cache regeneration.
#   Default: Qwen/Qwen2.5-7B-Instruct (fp16 ~14 GB, fits 24 GB GPU comfortably).
#   Alt:     Qwen/Qwen3-8B            (fp16 ~16 GB; native thinking mode).
# Set ENABLE_THINKING=True only with a Qwen3 family model; Qwen2.5 ignores it.
# Bump MAX_NEW_TOKENS to ~2048 when ENABLE_THINKING=True (CoT eats tokens).
MODEL_NAME = "Qwen/Qwen2.5-7B-Instruct"
ENABLE_THINKING = False
MAX_NEW_TOKENS = 512
SEED = 42

# %%
set_global_seeds(SEED)

# %% [markdown]
# ## 1. Infrastructure Detection
#
# Automatically detect GPU and Neo4j availability.

# %%
GPU_AVAILABLE = torch.cuda.is_available()
print(f"GPU available: {GPU_AVAILABLE}")
if GPU_AVAILABLE:
    print(f"GPU: {torch.cuda.get_device_name(0)}")
    print(f"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB")

# Neo4j connection
NEO4J_URI = os.getenv("NEO4J_URI", "bolt://localhost:7687")
NEO4J_USER = os.getenv("NEO4J_USER", "neo4j")
NEO4J_PASSWORD = os.getenv("NEO4J_PASSWORD", "password")

from neo4j import GraphDatabase

NEO4J_DRIVER = GraphDatabase.driver(NEO4J_URI, auth=(NEO4J_USER, NEO4J_PASSWORD))
NEO4J_DRIVER.verify_connectivity()
print(f"Neo4j connected: {NEO4J_URI}")

# %% [markdown]
# ## 2. Load LLM for Batch Extraction
#
# Load Qwen2.5-7B-Instruct for relationship extraction. The model stays in memory
# for efficient batch processing.

# %%
LLM_MODEL = None
LLM_TOKENIZER = None
BATCH_SIZE = LLM_BATCH_SIZE

_cache_path = get_chapter_dir(23) / "output" / "supply_chain_cache" / "extracted_triples.parquet"
_need_llm = RERUN_EXTRACTION or not _cache_path.exists()

if _need_llm and GPU_AVAILABLE:
    from transformers import AutoModelForCausalLM, AutoTokenizer

    try:
        print(f"Loading {MODEL_NAME} for batch extraction...")
        # MODEL_NAME comes from the parameters cell (overridable via Papermill).
        LLM_TOKENIZER = AutoTokenizer.from_pretrained(MODEL_NAME, padding_side="left")
        if LLM_TOKENIZER.pad_token is None:
            LLM_TOKENIZER.pad_token = LLM_TOKENIZER.eos_token

        LLM_MODEL = AutoModelForCausalLM.from_pretrained(
            MODEL_NAME,
            dtype=torch.float16,
            device_map="cuda",
        )
        print(f"Model loaded: {MODEL_NAME}")
        print(f"Batch size: {BATCH_SIZE}")
    except Exception as e:
        raise RuntimeError(f"Could not load {MODEL_NAME} on CUDA") from e
elif not _need_llm:
    print("Using the validated extraction cache; Qwen is not loaded")
else:
    raise RuntimeError("Cache regeneration requires a CUDA GPU")

# %% [markdown]
# ## 3. Data Loading
#
# Load pre-downloaded S&P 100 filings from the download script. The notebook now
# requires real staged filings instead of falling back to a demo dataset.

# %%
from data import load_sec_filings

# %%
filings_df = load_sec_filings("10-K", universe="sp100")
required_filing_columns = {
    "symbol",
    "company_name",
    "form",
    "filing_date",
    "accession_no",
    "year",
    "text",
    "text_length",
}
assert required_filing_columns <= set(filings_df.columns)
assert filings_df.filter(
    pl.any_horizontal(pl.col(list(required_filing_columns)).is_null())
).is_empty()
assert filings_df.select(pl.struct(["symbol", "accession_no"]).is_duplicated().sum()).item() == 0
assert filings_df.filter(pl.col("text").str.len_chars() != pl.col("text_length")).is_empty()
print(f"Loaded {len(filings_df)} 10-K filings via load_sec_filings()")

# Subsample for tractable Qwen runtime, but only when the LLM is the source of the
# triples. On the cached path the committed triples came from the whole corpus, so
# subsampling here would relabel every figure with a corpus that produced none of them.
if MAX_COMPANIES > 0 and RERUN_EXTRACTION:
    keep = sorted(filings_df["symbol"].unique().to_list())[:MAX_COMPANIES]
    filings_df = filings_df.filter(pl.col("symbol").is_in(keep))
    print(
        f"Subset to MAX_COMPANIES={MAX_COMPANIES}: {len(filings_df)} filings "
        f"across {filings_df['symbol'].n_unique()} companies"
    )
elif MAX_COMPANIES > 0:
    print(
        f"MAX_COMPANIES={MAX_COMPANIES} ignored: RERUN_EXTRACTION is False, and the "
        "cached triples were extracted from the full corpus"
    )

# Show summary
print(f"\nCompanies: {filings_df['symbol'].n_unique()}")
print(f"Filings: {len(filings_df)}")
if "year" in filings_df.columns:
    years = filings_df["year"].unique().sort().to_list()
    print(f"Years: {years}")

# %% [markdown]
# ## 4. Relationship Schema
#
# Define the knowledge graph schema. Each extracted relationship is a
# subject-predicate-object triple restricted to three financial relationship types.

# %% [markdown]
# ### Triple Dataclass
#
# Lightweight container for a single knowledge graph edge. The `to_dict()` method
# enables batch serialization for Neo4j UNWIND loading.


# %%
@dataclass
class Triple:
    """A subject-predicate-object relationship."""

    subject: str
    predicate: str
    object: str

    def to_dict(self) -> dict:
        return {"subject": self.subject, "predicate": self.predicate, "object": self.object}


# %%
RELATIONSHIP_TYPES = ["HAS_SUPPLIER", "COMPETES_WITH", "HAS_CUSTOMER"]

EXTRACTION_PROMPT = """You are a financial analyst extracting business relationships from SEC 10-K filings.

Extract ONLY the following relationship types:
1. HAS_SUPPLIER: Company depends on another entity for components/services
2. COMPETES_WITH: Company competes with another entity in a market
3. HAS_CUSTOMER: Company sells to another entity (if mentioned)

Output format: JSON array of objects with keys "subject", "predicate", "object"
- subject: The company name (use full official name)
- predicate: One of HAS_SUPPLIER, COMPETES_WITH, HAS_CUSTOMER
- object: The related entity name (abbreviated, e.g., "TSMC" not full name)

RULES:
- Extract ONLY explicitly stated relationships
- Entity names should be ≤5 words
- Normalize common names (TSMC, Foxconn, Samsung, etc.)

Example:
[{"subject": "Apple Inc.", "predicate": "HAS_SUPPLIER", "object": "TSMC"}]
"""

# %% [markdown]
# ## 5. Batch LLM Extraction
#
# Process multiple documents in parallel for higher throughput. The pipeline sends
# batches of filing texts through the LLM in a single forward pass, then parses
# the JSON output into structured triples.

# %% [markdown]
# ### Batch Extraction Function
#
# Core extraction function that sends multiple filing texts through the LLM in one
# batched forward pass.

# %% [markdown]
# ### Prompt Builder
#
# Construct a chat-formatted prompt for one filing so the batch extraction cell
# can stay focused on tokenization, generation, and JSON parsing.


# %%
def build_batch_prompt(text: str, company_name: str) -> str:
    """Create the chat-formatted extraction prompt for one filing."""
    user_prompt = f"""Company: {company_name}

SEC 10-K Filing Text:
{text[:6000]}

Extract all business relationships. Output ONLY valid JSON array."""
    messages = [
        {"role": "system", "content": EXTRACTION_PROMPT},
        {"role": "user", "content": user_prompt},
    ]
    # Qwen2.5 ignores enable_thinking; Qwen3 uses it when explicitly selected.
    return LLM_TOKENIZER.apply_chat_template(
        messages,
        tokenize=False,
        add_generation_prompt=True,
        enable_thinking=ENABLE_THINKING,
    )


# %% [markdown]
# ### Batch Extraction Function
#
# Send a batch of prompts through the LLM, decode the generated text, and parse
# each response into structured triples.


# %%
def extract_relationships_batch(texts: list[str], company_names: list[str]) -> list[list[Triple]]:
    """
    Extract relationships from multiple texts in a single batch.

    Uses batched generation for efficient GPU utilization.
    """
    if LLM_MODEL is None or LLM_TOKENIZER is None:
        raise RuntimeError(
            "GPU-backed Qwen2.5-7B extraction is required. The sample-triple "
            "fallback has been removed."
        )

    model_inputs = LLM_TOKENIZER(
        [
            build_batch_prompt(text, company_name)
            for text, company_name in zip(texts, company_names, strict=False)
        ],
        return_tensors="pt",
        padding=True,
        truncation=True,
        max_length=4096,
    ).to("cuda")

    with torch.no_grad():
        generated_ids = LLM_MODEL.generate(
            **model_inputs,
            max_new_tokens=MAX_NEW_TOKENS,
            do_sample=False,
            pad_token_id=LLM_TOKENIZER.pad_token_id,
        )

    # Decode responses
    results = []
    for i, (input_ids, output_ids) in enumerate(
        zip(model_inputs.input_ids, generated_ids, strict=False)
    ):
        # Get only generated tokens
        generated = output_ids[len(input_ids) :]
        response = LLM_TOKENIZER.decode(generated, skip_special_tokens=True)

        # Parse JSON
        triples = _parse_json_triples(response, company_names[i])
        results.append(triples)

    return results


# %% [markdown]
# ### JSON Response Parser
#
# Extract the JSON array from LLM output, handling common formatting issues
# (preamble text, trailing content). Only keeps triples with valid predicate types.


# %%
def _parse_json_triples(response: str, company_name: str) -> list[Triple]:
    """Parse JSON response into Triple objects."""
    # Strip Qwen3 thinking blocks if present (Qwen2.5 never emits them).
    response = re.sub(r"<think>.*?</think>", "", response, flags=re.DOTALL).strip()
    try:
        start_idx = response.find("[")
        end_idx = response.rfind("]") + 1
        if start_idx >= 0 and end_idx > start_idx:
            json_str = response[start_idx:end_idx]
            extracted = json.loads(json_str)
            return [
                Triple(
                    subject=item.get("subject", company_name),
                    predicate=item.get("predicate", ""),
                    object=item.get("object", ""),
                )
                for item in extracted
                if item.get("predicate") in RELATIONSHIP_TYPES
            ]
    except (json.JSONDecodeError, KeyError):
        pass
    return []


# %% [markdown]
# ### Batch Iterator Utility
#
# Simple chunking helper for processing filings in GPU-friendly batch sizes.


# %%
def batched(iterable, n: int) -> Generator:
    """Yield successive n-sized chunks from iterable."""
    items = list(iterable)
    for i in range(0, len(items), n):
        yield items[i : i + n]


# %% [markdown]
# ## 6. Extract Relationships
#
# The preserved producer run processed 601 filings through Qwen2.5-7B-Instruct
# in batches of 2. It took about 27 minutes on an NVIDIA RTX 3090 and used
# roughly 14 GB of VRAM. CPU regeneration is not a supported production path.
#
# **Pre-computed cache**: The repository ships a cached extraction result
# (`output/supply_chain_cache/extracted_triples.parquet`, 11KB) derived from
# public SEC EDGAR 10-K filings. Its Parquet bytes and sidecar metadata were
# committed with the full producer artifact. This lets you explore the graph analysis,
# Neo4j loading, and visualization sections without a GPU.
#
# To run the full LLM extraction yourself:
# ```python
# RERUN_EXTRACTION = True  # in the parameters cell above
# ```
#
# A sidecar `extracted_triples.meta.json` pins the parquet's content hash,
# schema, row count, and extractor identity. Reads recompute the hash and
# fail if the cache has drifted from the recorded producer output.

# %%
EXPECTED_CACHE_COLUMNS = ("subject", "predicate", "object")
CACHE_ROWS_MIN = 100
CACHE_ROWS_MAX = 50_000


def _hash_cache_bytes(path: Path) -> str:
    return hashlib.sha256(path.read_bytes()).hexdigest()


# %% [markdown]
# ### Cache Validator
#
# Validate the cache bytes, schema, and row count before constructing graph edges.


# %%
def _validate_cache(parquet_path: Path, meta_path: Path) -> pl.DataFrame:
    if not meta_path.exists():
        raise FileNotFoundError(
            f"Cache parquet exists but {meta_path.name} is missing. "
            "Delete the parquet or set RERUN_EXTRACTION=True to regenerate both."
        )
    meta = json.loads(meta_path.read_text())
    current_hash = _hash_cache_bytes(parquet_path)
    if current_hash != meta["content_hash"]:
        raise ValueError(
            f"Cache hash mismatch for {parquet_path.name}: "
            f"file={current_hash[:16]}, meta={meta['content_hash'][:16]}. "
            "Delete both files or set RERUN_EXTRACTION=True."
        )
    df = pl.read_parquet(parquet_path)
    if tuple(df.columns) != EXPECTED_CACHE_COLUMNS:
        raise ValueError(
            f"Cache schema drift: columns={df.columns}, expected={list(EXPECTED_CACHE_COLUMNS)}."
        )
    n_rows = df.height
    if not (CACHE_ROWS_MIN <= n_rows <= CACHE_ROWS_MAX):
        raise ValueError(
            f"Cache row count {n_rows} outside plausible range "
            f"[{CACHE_ROWS_MIN}, {CACHE_ROWS_MAX}]."
        )
    if n_rows != meta["row_count"]:
        raise ValueError(
            f"Cache row count {n_rows} disagrees with meta row_count {meta['row_count']}."
        )
    return df


# %% [markdown]
# ### Cache Metadata Writer
#
# Record the exact extraction identity whenever regeneration is explicitly enabled.


# %%
def _write_cache_meta(parquet_path: Path, meta_path: Path, n_rows: int) -> None:
    meta = {
        "content_hash": _hash_cache_bytes(parquet_path),
        "schema": list(EXPECTED_CACHE_COLUMNS),
        "row_count": n_rows,
        "model_name": MODEL_NAME,
        "max_companies": MAX_COMPANIES,
        "written_at": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
    }
    meta_path.write_text(json.dumps(meta, indent=2))


# %% [markdown]
# ### Extraction Runner
#
# Regenerate all candidate triples only when the parameter explicitly requests it.


# %%
def run_full_extraction(filing_records: list[dict]) -> tuple[list[Triple], float]:
    """Run deterministic batch extraction with visible progress."""
    started = time.time()
    extracted: list[Triple] = []
    n_batches = (len(filing_records) + BATCH_SIZE - 1) // BATCH_SIZE
    for batch_idx, batch in enumerate(batched(filing_records, BATCH_SIZE)):
        texts = [record["text"][:6000] for record in batch]
        names = [record["company_name"] for record in batch]
        batch_results = extract_relationships_batch(texts, names)
        for triples in batch_results:
            extracted.extend(triples)
        elapsed = time.time() - started
        rate = (batch_idx + 1) / elapsed if elapsed else 0.0
        remaining = (n_batches - batch_idx - 1) / rate if rate else 0.0
        print(
            f"Batch {batch_idx + 1}/{n_batches}: {len(batch)} filings, "
            f"{sum(map(len, batch_results))} relationships, {len(extracted)} total "
            f"[{elapsed:.0f}s elapsed, about {remaining:.0f}s remaining]",
            flush=True,
        )
        torch.cuda.empty_cache()
    return extracted, time.time() - started


# %%
CACHE_DIR = get_chapter_dir(23) / "output" / "supply_chain_cache"
CACHE_PATH = CACHE_DIR / "extracted_triples.parquet"
CACHE_META_PATH = CACHE_DIR / "extracted_triples.meta.json"

if not RERUN_EXTRACTION:
    cached_df = _validate_cache(CACHE_PATH, CACHE_META_PATH)
    CACHE_META = json.loads(CACHE_META_PATH.read_text())
    all_triples = [
        Triple(r["subject"], r["predicate"], r["object"]) for r in cached_df.iter_rows(named=True)
    ]
    extraction_elapsed = None
    # MODEL_NAME says what would run, not what did: the cache carries its own
    # producer, and overriding the parameter cannot retroactively change it.
    EXTRACTOR_NAME = CACHE_META["model_name"]
    CACHE_CONTENT_HASH = CACHE_META["content_hash"]
    print(f"Loaded and validated {len(all_triples)} cached triples from {CACHE_PATH.name}")
    print(f"Cache producer: {EXTRACTOR_NAME}, {CACHE_META['row_count']} rows")
    # Narrow to the cohort the producer recorded. A cache regenerated with
    # MAX_COMPANIES set covers part of the corpus, and the roster and figures
    # below have to describe the filings the triples came from.
    CACHE_SCOPE = int(CACHE_META.get("max_companies", 0) or 0)
    if CACHE_SCOPE:
        cached_symbols = sorted(filings_df["symbol"].unique().to_list())[:CACHE_SCOPE]
        filings_df = filings_df.filter(pl.col("symbol").is_in(cached_symbols))
        years = filings_df["year"].unique().sort().to_list()
        print(
            f"Cache covers MAX_COMPANIES={CACHE_SCOPE}: narrowed to {len(filings_df)} "
            f"filings across {filings_df['symbol'].n_unique()} companies"
        )
else:
    all_triples, extraction_elapsed = run_full_extraction(filings_df.to_dicts())
    cache_df = pl.DataFrame([triple.to_dict() for triple in all_triples])
    cache_df.write_parquet(CACHE_PATH)
    _write_cache_meta(CACHE_PATH, CACHE_META_PATH, len(all_triples))
    EXTRACTOR_NAME = MODEL_NAME
    CACHE_CONTENT_HASH = _hash_cache_bytes(CACHE_PATH)
    print(f"Regenerated {len(all_triples)} triples in {extraction_elapsed:.1f}s")

# %% [markdown]
# ## 7. Entity Resolution
#
# The extraction prompt hands the model the filer's name and asks it to repeat that
# name as the subject. It comes back unchanged for 51 of the 133 distinct subject
# strings in the cache. Another 61 are the same registrant re-cased, re-punctuated or
# word-order swapped, and 21 name something else, usually an operating subsidiary or
# a brand. Each spelling is a separate node unless something merges them, and every
# degree count downstream is then measured on the fragments.
#
# Two things resolve names here. A canonical key strips case, punctuation, corporate
# suffixes and word order, so all spellings of one name collapse onto one node. The
# key is then looked up in the roster of S&P 100 filers built from the corpus itself,
# and a match takes the filer's official name. That second step is what makes a
# company named as somebody's competitor the same node as the company that filed its
# own 10-K.

# %% [markdown]
# ### Canonical Entity Key
#
# The join key for every name in the graph. Sorting the surviving tokens is what
# merges `SCHWAB CHARLES CORP` with `The Charles Schwab Corporation`; it is the only
# rule here that can merge names with different word order, and on this cache it
# merges exactly one group, the four spellings of that company.

# %%
CORPORATE_SUFFIXES = frozenset(
    {
        "inc",
        "incorporated",
        "corp",
        "corporation",
        "co",
        "cos",
        "company",
        "companies",
        "plc",
        "ltd",
        "limited",
        "lp",
        "llc",
        "holdings",
        "holding",
        "group",
        "the",
        "nv",
        "sa",
        "ag",
        "new",
        "del",
        "de",
    }
)


def entity_key(name: str) -> str:
    """Canonical join key: case, punctuation, corporate suffix and word order removed."""
    stripped = re.sub(r"[^a-z0-9]+", " ", name.lower())
    tokens = [token for token in stripped.split() if token not in CORPORATE_SUFFIXES]
    return " ".join(sorted(tokens or stripped.split()))


def normalize_entity_text(name: str) -> str:
    """Collapse whitespace and strip trailing punctuation from extracted names."""
    return " ".join(name.split()).strip(" ,.;:")


# %% [markdown]
# ### Filer Roster
#
# The corpus knows the official name of every company whose filings were extracted,
# so the subject of a triple is not a name that has to be guessed at. Alphabet files
# under two symbols with one name, which is why the roster is keyed on the name.

# %%
FILER_ROSTER: dict[str, str] = {}
for filer_name in sorted(set(filings_df["company_name"].to_list())):
    FILER_ROSTER.setdefault(entity_key(filer_name), filer_name)
assert len(FILER_ROSTER) == len(set(filings_df["company_name"].to_list())), (
    "two filers share a canonical key; the roster cannot resolve between them"
)
print(f"Filer roster: {len(FILER_ROSTER)} issuers from {filings_df['symbol'].n_unique()} symbols")

# %% [markdown]
# ### Entity Alias Map
#
# Short forms for entities that never file with the SEC, so the roster cannot reach
# them. The extraction prompt already tells the model to emit `TSMC` rather than the
# full legal name, and the hit counts printed below show what that leaves for the
# alias map to do.

# %%
ENTITY_ALIASES = {
    "Taiwan Semiconductor Manufacturing Company": "TSMC",
    "Taiwan Semiconductor": "TSMC",
    "Foxconn Technology Group": "Foxconn",
    "Hon Hai Precision": "Foxconn",
    "Samsung Electronics": "Samsung",
    "SK Hynix": "SK Hynix",
    "Hynix": "SK Hynix",
    "Advanced Micro Devices": "AMD",
    "Amazon Web Services": "AWS",
    "Microsoft Azure": "Azure",
}
ALIAS_BY_KEY = {entity_key(variant): standard for variant, standard in ENTITY_ALIASES.items()}
assert len(ALIAS_BY_KEY) == len(ENTITY_ALIASES), "two aliases collapse onto one canonical key"

# %% [markdown]
# ### Generic Entity Filter
#
# Reject category phrases that do not identify a company or organization. The filter
# matches on the canonical key rather than the raw string, so a phrase cannot escape
# it by arriving in a different case or with a hyphen the list does not carry.

# %%
# Generic category phrases are not named graph entities. A pipe-delimited string
# keeps this reader-facing configuration compact.
GENERIC_ENTITY_NAMES = """
aluminum suppliers|broadcast station group|broadcast station groups|composite suppliers
media corporation|media corporations|numerous suppliers|restaurants|third party suppliers
unaffiliated third party suppliers|consumers|third parties|third party|third-party
third-party manufacturer|third-party manufacturers|third-party tower operators
financial institutions|financial institution|governments|government|government agencies
distributors|wholesale distributors|vendors|wholesalers|travelers|hospitals|merchants
leaf merchants|domestic tobacco growers|fintechs|businesses|pharmaceutical companies
manufacturers|sensor manufacturers|electronics|electronics manufacturing service providers
travel service providers|enterprises|banks|airlines|routers|infrastructure equipment
insurance company clients|academic|energy|healthcare|semiconductor|hd mapping companies
sellers|buyers|startups|suppliers|research and industrial|food & beverage|power & renewables
agricultural|grapes|patients and communities|business and general aviation aircraft operators
mobile providers|other multichannel video providers|digital messaging and payment platforms
high-frequency stores|3000+ small businesses|small, minority- and women-owned businesses
construction, earthmoving, material handling, roadbuilding, and/or forestry equipment locations
""".strip()
GENERIC_ENTITY_KEYS = frozenset(
    entity_key(name.strip()) for name in GENERIC_ENTITY_NAMES.replace("\n", "|").split("|")
)
assert not (GENERIC_ENTITY_KEYS & set(FILER_ROSTER)), "a generic phrase collides with a filer"


# %%
def is_actionable_entity(name: str, predicate: str) -> bool:
    """Reject generic placeholders that do not identify a real graph node."""
    lower = name.lower()
    if not name or entity_key(name) in GENERIC_ENTITY_KEYS:
        return False
    if predicate == "HAS_SUPPLIER" and "supplier" in lower and lower not in {"supplier.io"}:
        return False
    return not (predicate == "HAS_CUSTOMER" and "customer" in lower)


# %% [markdown]
# ### Canonical Name Chooser
#
# Every name sharing a key becomes one node, so one surface form has to represent the
# group. A filer match takes the corpus spelling. Otherwise the most frequent form
# wins, with all-caps forms broken to the back of the tie, so a group holding one
# `Apple` and one `APPLE` resolves to `Apple`.


# %%
def build_canonical_names(triples: list[Triple]) -> tuple[dict[str, str], dict[str, int]]:
    """Map every extracted name to one canonical form, and count the alias rewrites."""
    frequency: Counter[str] = Counter()
    forms: dict[str, set[str]] = {}
    for triple in triples:
        for raw in (triple.subject, triple.object):
            name = normalize_entity_text(raw)
            if not name:
                continue
            frequency[name] += 1
            forms.setdefault(entity_key(name), set()).add(name)

    canonical: dict[str, str] = {}
    alias_hits: Counter[str] = Counter()
    for key, group in forms.items():
        if key in FILER_ROSTER:
            chosen = FILER_ROSTER[key]
        elif key in ALIAS_BY_KEY:
            chosen = ALIAS_BY_KEY[key]
            alias_hits[key] += sum(frequency[name] for name in group)
        else:
            chosen = sorted(group, key=lambda name: (-frequency[name], name.isupper(), name))[0]
        for name in group:
            canonical[name] = chosen
    return canonical, alias_hits


def resolve_entity(name: str, canonical: dict[str, str]) -> str:
    """Resolve one extracted name to its canonical graph node name."""
    cleaned = normalize_entity_text(name)
    return canonical.get(cleaned, cleaned)


# %%
CANONICAL_NAMES, ALIAS_HITS = build_canonical_names(all_triples)

resolved_triples = []
for t in all_triples:
    resolved_subject = resolve_entity(t.subject, CANONICAL_NAMES)
    resolved_object = resolve_entity(t.object, CANONICAL_NAMES)
    if not is_actionable_entity(resolved_subject, t.predicate):
        continue
    if not is_actionable_entity(resolved_object, t.predicate):
        continue
    resolved_triples.append(
        Triple(subject=resolved_subject, predicate=t.predicate, object=resolved_object)
    )

# Deduplicate
unique_triples = list({(t.subject, t.predicate, t.object): t for t in resolved_triples}.values())
print(f"After resolution and dedup: {len(unique_triples)} unique triples")

# %% [markdown]
# ### What Resolution Actually Did
#
# Print the work rather than asserting it happened. The alias table shows which
# entries rewrote a name and which were unreachable, either because the roster
# already covers that entity or because no spelling in the cache matches them.

# %%
raw_names = {normalize_entity_text(n) for t in all_triples for n in (t.subject, t.object) if n}
merged_names = {CANONICAL_NAMES[n] for n in raw_names if n in CANONICAL_NAMES}
raw_subjects = {normalize_entity_text(t.subject) for t in all_triples}
matched_subjects = {n for n in raw_subjects if entity_key(n) in FILER_ROSTER}
print(f"Distinct extracted names: {len(raw_names)} -> {len(merged_names)} after resolution")
print(
    f"Subject strings matching a filer: {len(matched_subjects)}/{len(raw_subjects)} "
    f"({len(raw_subjects - matched_subjects)} name entities the roster cannot reach)"
)
print("\nAlias map, by rewritten name slots:")
for variant, standard in ENTITY_ALIASES.items():
    key = entity_key(variant)
    if key in FILER_ROSTER:
        reason = f"unused: the filer roster resolves it to {FILER_ROSTER[key]}"
    elif ALIAS_HITS[key]:
        slots = ALIAS_HITS[key]
        reason = f"rewrote {slots} name slot{'s' if slots != 1 else ''}"
    else:
        reason = "unused: no spelling in the cache matches this variant"
    print(f"  {variant} -> {standard}: {reason}")

# %% [markdown]
# ### The Names the Roster Cannot Reach
#
# A subject that matches no filer is not noise. Most of them are operating
# subsidiaries and brands that the model named instead of the registrant, and no
# string rule can fold `KAYAK` into `Booking Holdings Inc.` A production pipeline
# resolves these against a corporate hierarchy; this one leaves them as their own
# nodes, so the company count below is larger than the number of filers in it.

# %%
unreachable_subjects = sorted(
    {t.subject for t in unique_triples if entity_key(t.subject) not in FILER_ROSTER}
)
print(f"Subject nodes that are not S&P 100 filers: {len(unreachable_subjects)}")
for name in unreachable_subjects:
    print(f"  {name}")

# %% [markdown]
# ## 8. Graph Statistics
#
# Analyze the extracted knowledge graph to identify concentration risk and
# shared dependencies.

# %%
# Compute statistics
subjects = set(t.subject for t in unique_triples)
objects = set(t.object for t in unique_triples)
all_entities = subjects | objects
companies = subjects
suppliers = {t.object for t in unique_triples if t.predicate == "HAS_SUPPLIER"}
competitors = {t.object for t in unique_triples if t.predicate == "COMPETES_WITH"}
customers = {t.object for t in unique_triples if t.predicate == "HAS_CUSTOMER"}

filer_companies = {c for c in companies if entity_key(c) in FILER_ROSTER}

supplier_rels = sum(1 for t in unique_triples if t.predicate == "HAS_SUPPLIER")
competitor_rels = sum(1 for t in unique_triples if t.predicate == "COMPETES_WITH")
customer_rels = sum(1 for t in unique_triples if t.predicate == "HAS_CUSTOMER")

# %%
# Find shared suppliers (critical concentration risk nodes)
supplier_companies = {}
for t in unique_triples:
    if t.predicate == "HAS_SUPPLIER":
        if t.object not in supplier_companies:
            supplier_companies[t.object] = set()
        supplier_companies[t.object].add(t.subject)

shared_suppliers = {s: cs for s, cs in supplier_companies.items() if len(cs) > 1}
assert shared_suppliers, (
    "no supplier is named by more than one company, so the shared-supplier figures "
    "below have nothing to draw. A subsampled extraction run reaches this state."
)

# %%
print("=" * 60)
print("KNOWLEDGE GRAPH STATISTICS")
print("=" * 60)
print(f"Companies analyzed:       {len(companies)}")
print(f"Unique suppliers:         {len(suppliers)}")
print(f"Unique competitors:       {len(competitors)}")
print(f"Unique customers:         {len(customers)}")
print(f"Total entities:           {len(all_entities)}")
print(f"Supplier relationships:   {supplier_rels}")
print(f"Competitor relationships: {competitor_rels}")
print(f"Customer relationships:   {customer_rels}")
print(f"Total relationships:      {len(unique_triples)}")
print(f"  of the companies, S&P 100 filers: {len(filer_companies)}")
print(f"  suppliers named by one company:   {len(suppliers) - len(shared_suppliers)}")
print(f"  suppliers named by more than one: {len(shared_suppliers)}")

if shared_suppliers:
    print("\nSuppliers named by more than one company:")
    sorted_shared = sorted(shared_suppliers.items(), key=lambda x: len(x[1]), reverse=True)
    for supplier, companies_set in sorted_shared[:10]:
        print(f"  {supplier}: {len(companies_set)} companies")

# %% [markdown]
# Almost every extracted supplier is named by exactly one company, so the shared
# structure this chapter queries rests on the handful above. That is a property of
# what a 10-K names rather than of the supply chain: a filer lists the suppliers it
# is required to disclose, and only the ones many filers depend on get repeated.
#
# Shared-supplier degree measures exposure concentration in the extracted graph. It
# does not establish disruption probabilities or causal propagation. The metric is
# useful as a portfolio diagnostic after the extracted entities have been reviewed.

# %% [markdown]
# ### Relationship Type Distribution
#
# The edge classes, and the supplier degree distribution behind the shared-supplier
# claim. Panel (b) plots every supplier named by more than one company, so its bar
# count is the finding rather than a top-N cut.

# %%
RELATIONSHIP_FIGURE_ALT = (
    f"Two stacked panels. The upper bar chart counts extracted edges by class: "
    f"{competitor_rels} competitor, {customer_rels} customer and {supplier_rels} "
    f"supplier relationships. The lower horizontal bar chart shows the "
    f"{len(shared_suppliers)} suppliers named by more than one company: "
    + ", ".join(
        f"{name} at {len(cs)}"
        for name, cs in sorted(shared_suppliers.items(), key=lambda kv: -len(kv[1]))
    )
    + " companies."
)

fig, axes = plt.subplots(2, 1, figsize=FIGSIZE["dual_v"], constrained_layout=True)

# Panel (a): Relationship counts
rel_types = ["Supplier", "Competitor", "Customer"]
rel_counts = [supplier_rels, competitor_rels, customer_rels]
bars = axes[0].bar(rel_types, rel_counts, color=[COLORS["blue"], COLORS["amber"], COLORS["copper"]])
axes[0].set_ylabel("Number of Relationships")
for bar, count in zip(bars, rel_counts):
    axes[0].text(
        bar.get_x() + bar.get_width() / 2,
        bar.get_height() + 3,
        str(count),
        ha="center",
        fontweight="bold",
    )

# Panel (b): every supplier named by more than one company, not a top-N slice
sorted_shared = sorted(shared_suppliers.items(), key=lambda x: len(x[1]), reverse=True)
names = [s[:20] for s, _ in sorted_shared]
counts = [len(cs) for _, cs in sorted_shared]
axes[1].barh(range(len(names)), counts, color=COLORS["amber"])
axes[1].set_yticks(range(len(names)))
axes[1].set_yticklabels(names, fontsize=8)
axes[1].set_xlabel("Companies naming this supplier")
axes[1].set_title("Suppliers named by more than one company", loc="left")
axes[1].invert_yaxis()
axes[1].text(
    0.99,
    0.05,
    f"{len(suppliers) - len(shared_suppliers)} of {len(suppliers)} extracted suppliers\n"
    "are named by one company only",
    transform=axes[1].transAxes,
    ha="right",
    va="bottom",
    fontsize=8,
    color=COLORS["slate"],
)

add_message_title(
    axes[0],
    "Extracted relationship classes and their shared suppliers",
    subtitle=(
        f"{EXTRACTOR_NAME} extraction from the S&P 100 10-K filings, {min(years)}-{max(years)}"
    ),
)
show_with_alt(fig, RELATIONSHIP_FIGURE_ALT)

# %% [markdown]
# ## 9. Neo4j Batch Loading
#
# Use UNWIND for efficient batch graph loading. UNWIND takes a list parameter and
# expands it into rows, allowing a single Cypher statement to create hundreds of
# nodes and relationships in one transaction instead of one statement per relationship.

# %% [markdown]
# ### Batch Loader
#
# Loads all triples to Neo4j in two passes (suppliers, then competitors), using
# MERGE to avoid duplicates while only clearing the supply-chain relationships
# created by this notebook.

# %% [markdown]
# ### Batch Query Runner
#
# Execute one UNWIND statement for a homogeneous triple batch and return the
# number of relationships loaded.


# %%
def run_unwind_batch(session, batch: list[Triple], query: str) -> int:
    """Execute one UNWIND load for a batch of triples."""
    batch_data = [triple.to_dict() for triple in batch]
    session.run(query, batch=batch_data)
    return len(batch)


# %% [markdown]
# ### UNWIND Templates per Predicate
#
# Each supply-chain predicate has its own MERGE template. Every template merges on
# `:Entity {name}` and adds the role as a second label, because a name is one entity
# no matter which side of an edge it appears on. Merging on the role label instead
# gives Microsoft three separate nodes here: it is a company that files, a supplier
# somebody names, and a customer somebody else names.


# %%
UNWIND_TEMPLATES: dict[str, str] = {
    "HAS_SUPPLIER": """
        UNWIND $batch AS row
        MERGE (c:Entity {name: row.subject})
        SET c:Company
        MERGE (s:Entity {name: row.object})
        SET s:Supplier
        MERGE (c)-[:HAS_SUPPLIER]->(s)
        """,
    "COMPETES_WITH": """
        UNWIND $batch AS row
        MERGE (c1:Entity {name: row.subject})
        SET c1:Company
        MERGE (c2:Entity {name: row.object})
        SET c2:Company
        MERGE (c1)-[:COMPETES_WITH]->(c2)
        """,
    "HAS_CUSTOMER": """
        UNWIND $batch AS row
        MERGE (c:Entity {name: row.subject})
        SET c:Company
        MERGE (cust:Entity {name: row.object})
        SET cust:Customer
        MERGE (c)-[:HAS_CUSTOMER]->(cust)
        """,
}


# %% [markdown]
# ### Predicate-Scoped Loader
#
# Run `run_unwind_batch` across every batch of a single predicate's triples.


# %%
def _load_predicate(session, triples: list[Triple], predicate: str, batch_size: int) -> int:
    """Load all triples of a single predicate using its UNWIND template."""
    subset = [t for t in triples if t.predicate == predicate]
    loaded = 0
    for batch in batched(subset, batch_size):
        loaded += run_unwind_batch(session, batch, UNWIND_TEMPLATES[predicate])
    return loaded


# %% [markdown]
# ### Clearing an Existing Graph
#
# This notebook shares a database with the rest of the chapter, and it is the
# second writer to touch `:Company`: `08_8k_event_extraction` attaches its event
# relationships to nodes with that label and declares `Company.name` unique. So the
# reset cannot delete company nodes, and it cannot leave a second node behind for a
# name that already exists.
#
# Adopting existing `:Company` nodes into `:Entity` is what prevents the duplicate.
# Without it, a database where 08 ran first already holds `(:Company {name: "Apple
# Inc."})`, `MERGE (:Entity {name: "Apple Inc."})` matches nothing, and the load
# creates a second node that the uniqueness constraint then rejects.
#
# Deleting only entity nodes that have no relationships left is what protects 08:
# a `DETACH DELETE` over `:Entity` would take its `APPOINTED` and `ACQUIRED` edges
# with it, because after adoption those nodes carry the label too.


# %%
RESET_STATEMENTS = (
    # This notebook's own edges, whichever revision wrote them.
    "MATCH (:Company)-[r:HAS_SUPPLIER|COMPETES_WITH|HAS_CUSTOMER]->() DELETE r",
    # One node per name across the chapter: adopt company nodes written by 08, or
    # by an earlier revision of this notebook, before anything merges on :Entity.
    "MATCH (n:Company) WHERE NOT n:Entity SET n:Entity",
    # Role nodes from the revision that keyed on the role label are unreachable now.
    "MATCH (n:Supplier) WHERE NOT n:Entity DETACH DELETE n",
    "MATCH (n:Customer) WHERE NOT n:Entity DETACH DELETE n",
    # Roles are re-applied by the load; a stale one would outlive its edges.
    "MATCH (n:Entity) REMOVE n:Supplier, n:Customer",
    # What is left with no edges at all belongs to no notebook.
    "MATCH (n:Entity) WHERE NOT (n)--() DELETE n",
)


# %% [markdown]
# ### Top-Level Batch Loader
#
# Reset the supply-chain subgraph, then dispatch each predicate to `_load_predicate`.


# %%
def load_to_neo4j_batch(triples: list[Triple], batch_size: int = 1000) -> int:
    """
    Load triples to Neo4j using UNWIND for batch efficiency.

    Returns count of loaded triples.
    """
    if NEO4J_DRIVER is None:
        raise RuntimeError(
            "Neo4j is required for this load step. Start Neo4j and set "
            "NEO4J_URI, NEO4J_USER, and NEO4J_PASSWORD."
        )

    loaded = 0
    with NEO4J_DRIVER.session() as session:
        for statement in RESET_STATEMENTS:
            session.run(statement)
        for predicate in UNWIND_TEMPLATES:
            loaded += _load_predicate(session, triples, predicate, batch_size)

    print(f"Loaded {loaded} triples to Neo4j (batch size: {batch_size})")
    return loaded


print("\n" + "=" * 60)
print("NEO4J BATCH LOADING")
print("=" * 60)

load_start_time = time.time()
loaded_count = load_to_neo4j_batch(unique_triples, batch_size=500)
with NEO4J_DRIVER.session() as session:
    persisted_count = session.run(
        "MATCH (:Company)-[r:HAS_SUPPLIER|COMPETES_WITH|HAS_CUSTOMER]->() RETURN count(r) AS n"
    ).single()["n"]
    persisted_entities = session.run(
        "MATCH (n:Entity) WHERE EXISTS { (n)-[:HAS_SUPPLIER|COMPETES_WITH|HAS_CUSTOMER]-() } "
        "RETURN count(n) AS n"
    ).single()["n"]
assert persisted_count == len(unique_triples), (
    f"Neo4j contains {persisted_count} supply-chain edges; expected {len(unique_triples)}"
)
assert persisted_entities == len(all_entities), (
    f"Neo4j holds {persisted_entities} entity nodes on the supply-chain edges; "
    f"expected {len(all_entities)}. A name that reached the graph twice means "
    "resolution did not collapse it."
)
load_elapsed = time.time() - load_start_time
print(f"Loading completed in {load_elapsed:.2f}s")
print(f"Entity nodes on supply-chain edges: {persisted_entities}, edges: {persisted_count}")

# %% [markdown]
# ### Graph Snapshot
#
# `09_knowledge_graph_features` reads this graph back and needs to know it is reading
# what this notebook wrote, not a graph a partial run left behind. Record the identity
# here, computed from the graph as Neo4j returns it, so the check downstream compares
# two independent reads rather than a constant somebody remembered to update.
#
# The relationship query text is repeated in 09. If either copy drifts, the two reads
# order differently and the hash comparison fails, which is the intended outcome.

# %%
SNAPSHOT_NAME = "ch23_supply_chain"
RELATIONSHIP_QUERY = """
    CALL () {
        MATCH (c:Company)-[:HAS_SUPPLIER]->(s:Supplier)
        RETURN c.name AS company, 'HAS_SUPPLIER' AS predicate, s.name AS related
        UNION ALL
        MATCH (c:Company)-[:COMPETES_WITH]->(peer:Company)
        RETURN c.name AS company, 'COMPETES_WITH' AS predicate, peer.name AS related
        UNION ALL
        MATCH (c:Company)-[:HAS_CUSTOMER]->(cust:Customer)
        RETURN c.name AS company, 'HAS_CUSTOMER' AS predicate, cust.name AS related
    }
    RETURN company, predicate, related
    ORDER BY company, predicate, related
"""


def read_graph_identity(session) -> tuple[str, dict[str, int], int]:
    """Hash the supply-chain edges as Neo4j returns them, with per-class counts."""
    lines, counts = [], Counter()
    for record in session.run(RELATIONSHIP_QUERY):
        company = " ".join(record["company"].split())
        related = " ".join(record["related"].split())
        lines.append(f"{company}\t{record['predicate']}\t{related}")
        counts[record["predicate"]] += 1
    digest = hashlib.sha256("\n".join(lines).encode()).hexdigest()
    return digest, dict(counts), len(lines)


with NEO4J_DRIVER.session() as session:
    graph_sha256, edge_counts, edge_total = read_graph_identity(session)
    SUPPLY_SNAPSHOT = {
        "name": SNAPSHOT_NAME,
        "graph_sha256": graph_sha256,
        "company_count": len(companies),
        "entity_count": persisted_entities,
        "edge_total": edge_total,
        "supplier_edges": edge_counts.get("HAS_SUPPLIER", 0),
        "competitor_edges": edge_counts.get("COMPETES_WITH", 0),
        "customer_edges": edge_counts.get("HAS_CUSTOMER", 0),
        "cache_content_hash": CACHE_CONTENT_HASH,
        "extractor": EXTRACTOR_NAME,
    }
    session.run("MERGE (s:SupplyGraphSnapshot {name: $s.name}) SET s = $s", s=SUPPLY_SNAPSHOT)

assert edge_total == len(unique_triples), (
    f"the snapshot read {edge_total} edges back; the loader wrote {len(unique_triples)}"
)
assert edge_counts == {
    "HAS_SUPPLIER": supplier_rels,
    "COMPETES_WITH": competitor_rels,
    "HAS_CUSTOMER": customer_rels,
}, f"edge classes read back as {edge_counts}"
print(f"Supply graph snapshot {SNAPSHOT_NAME}: sha256 {graph_sha256[:12]}, {edge_total} edges")

# %% [markdown]
# ## 10. Example Queries
#
# Cypher queries for supply chain analysis.

# %%
print("\n" + "=" * 60)
print("EXAMPLE CYPHER QUERIES")
print("=" * 60)

queries = {
    "Find all suppliers for a company": """
MATCH (c:Company {name: 'Apple Inc.'})-[:HAS_SUPPLIER]->(s:Supplier)
RETURN s.name AS Supplier ORDER BY s.name
""",
    "Find shared suppliers across multiple companies": """
MATCH (s:Supplier)<-[:HAS_SUPPLIER]-(c:Company)
WITH s, COLLECT(c.name) AS companies, COUNT(c) AS company_count
WHERE company_count > 1
RETURN s.name AS Supplier, company_count, companies
ORDER BY company_count DESC LIMIT 10
""",
    "Find competitor clusters": """
MATCH (c1:Company)-[:COMPETES_WITH]->(c2:Company)
RETURN c1.name AS Company, COLLECT(c2.name) AS Competitors
ORDER BY SIZE(COLLECT(c2.name)) DESC
""",
    "Supply chain risk - single points of failure": """
MATCH (s:Supplier)<-[:HAS_SUPPLIER]-(c:Company)
WITH s, COUNT(c) AS customer_count
WHERE customer_count >= 3
RETURN s.name AS CriticalSupplier, customer_count
ORDER BY customer_count DESC
""",
}

for name, query in queries.items():
    print(f"\n-- {name} --")
    print(query.strip())

# %% [markdown]
# ## 11. Summary Statistics

# %% [markdown]
# The two stages are timed separately and never summed. On the cached path the
# extraction did not run in this process at all, so a "total pipeline time" would be
# the Neo4j load wearing the label of the 27 GPU-minutes that produced the triples.

# %%
summary_stats = {
    "Metric": [
        "Companies analyzed",
        "  of which S&P 100 filers",
        "Entity nodes",
        "Total relationships",
        "Supplier relationships",
        "Competitor relationships",
        "Customer relationships",
        "Unique suppliers",
        "Suppliers named by 2+ companies",
        "Suppliers named by 3+ companies",
        "Extraction time (s)",
        "Neo4j load time (s)",
        "Neo4j edges per second",
    ],
    "Value": [
        str(len(companies)),
        str(len(filer_companies)),
        str(persisted_entities),
        str(len(unique_triples)),
        str(supplier_rels),
        str(competitor_rels),
        str(customer_rels),
        str(len(suppliers)),
        str(len(shared_suppliers)),
        str(len([s for s, cs in shared_suppliers.items() if len(cs) >= 3])),
        "not run (cached)" if extraction_elapsed is None else f"{extraction_elapsed:.1f}",
        f"{load_elapsed:.1f}",
        f"{len(unique_triples) / max(load_elapsed, 1e-6):.0f}",
    ],
}

summary_df = pl.DataFrame(summary_stats)
print("\n" + "=" * 60)
print("FINAL STATISTICS")
print("=" * 60)
print(summary_df)

# %% [markdown]
# ## 12. Network Visualizations
#
# Three visualization formats:
# 1. **Static (Book)**: Publication-ready matplotlib figure
# 2. **Interactive (Notebook)**: pyvis for exploration
# 3. **Web Export (D3)**: JSON data for website integration
#
# ### Trading Applications
#
# The network view describes shared suppliers and competitor clusters. It does not
# estimate disruption probabilities or returns; those require separate evidence.

# %%
import networkx as nx

# %% [markdown]
# ### Build NetworkX Graph
#
# Filter to the most connected suppliers and build a NetworkX graph for
# visualization. This focuses the network on the supply chain nodes with
# highest concentration risk.


# %%
def build_networkx_graph(triples: list[Triple], min_companies: int = 2) -> nx.Graph:
    """Build the subgraph of suppliers named by at least `min_companies` companies."""
    G = nx.Graph()

    named_by = {}
    for t in triples:
        if t.predicate == "HAS_SUPPLIER":
            named_by.setdefault(t.object, set()).add(t.subject)

    # A count cut, not a rank cut. Ranking and keeping the top N would fill the
    # figure with arbitrary members of the tie at one company, which is where all
    # but a handful of the extracted suppliers sit.
    kept = {s for s, cs in named_by.items() if len(cs) >= min_companies}

    for t in triples:
        if t.predicate == "HAS_SUPPLIER" and t.object in kept:
            G.add_node(t.subject, node_type="company", label=t.subject[:15])
            G.add_node(t.object, node_type="supplier", count=len(named_by[t.object]))
            G.add_edge(t.subject, t.object, edge_type="supplies")

    return G


# %%
# Build the graph
MIN_SHARED_COMPANIES = 2
G = build_networkx_graph(unique_triples, min_companies=MIN_SHARED_COMPANIES)
network_suppliers = [n for n, d in G.nodes(data=True) if d.get("node_type") == "supplier"]
network_companies = [n for n, d in G.nodes(data=True) if d.get("node_type") == "company"]
print(f"Network graph: {G.number_of_nodes()} nodes, {G.number_of_edges()} edges")
print(
    f"{len(network_suppliers)} suppliers named by {MIN_SHARED_COMPANIES}+ companies, "
    f"reaching {len(network_companies)} of the {len(companies)} companies in the graph"
)
assert len(network_suppliers) == len(shared_suppliers), (
    "the network figure and the shared-supplier count disagree"
)

# %% [markdown]
# ### 12.1 Static Book Figure
#
# Publication-ready matplotlib figure showing supply chain concentration risk.
# Suppliers are positioned in an inner ring (sized by connection count), companies
# in an outer ring. The most critical supplier is annotated.

# %%
import math

from matplotlib.lines import Line2D

# The shared notebook style is initialized by ``utils.style``.

# %% [markdown]
# ### Static Figure Function
#
# Builds a two-ring layout: suppliers in the inner circle (sized by degree),
# companies in the outer ring. Annotates the highest-degree supplier with a
# callout showing its dependency count.

# %% [markdown]
# ### Static Layout Builder
#
# Compute the ring layout once so the plotting function only handles rendering.


# %%
def build_static_layout(
    G: nx.Graph,
) -> tuple[list[str], list[str], dict[str, tuple[float, float]], list[str]]:
    """Separate node types and place them on two concentric rings."""
    suppliers = [n for n, d in G.nodes(data=True) if d.get("node_type") == "supplier"]
    companies = [n for n, d in G.nodes(data=True) if d.get("node_type") == "company"]
    pos: dict[str, tuple[float, float]] = {}
    sorted_suppliers = sorted(suppliers, key=lambda x: G.degree(x), reverse=True)
    for i, supplier in enumerate(sorted_suppliers):
        angle = 2 * math.pi * i / len(sorted_suppliers)
        pos[supplier] = (math.cos(angle), math.sin(angle))
    for i, company in enumerate(companies):
        angle = 2 * math.pi * i / len(companies)
        pos[company] = (2.5 * math.cos(angle), 2.5 * math.sin(angle))
    return suppliers, companies, pos, sorted_suppliers


# %% [markdown]
# ### Static Annotation Helper
#
# Keep the title, legend, and concentration callout in one helper so the main
# plotting cell stays concise.

# %% [markdown]
# ### Static Supplier Callout
#
# Highlight the most connected supplier so the figure immediately communicates
# concentration risk.


# %%
def add_static_supplier_callout(
    ax: plt.Axes, G: nx.Graph, pos: dict[str, tuple[float, float]], sorted_suppliers: list[str]
) -> None:
    """Annotate the most connected supplier in the static figure."""
    if not sorted_suppliers:
        return
    top = sorted_suppliers[0]
    top_pos = pos[top]
    ax.annotate(
        f"{G.degree(top)} companies\nname this supplier",
        xy=top_pos,
        xytext=(top_pos[0] + 0.8, top_pos[1] + 0.8),
        fontsize=9,
        ha="left",
        arrowprops={"arrowstyle": "->", "color": COLORS["copper"]},
        bbox={"boxstyle": "round,pad=0.3", "facecolor": "white", "edgecolor": COLORS["copper"]},
    )


# %% [markdown]
# ### Static Legend
#
# Keep the legend handles in a constant so the figure finalizer stays short.

# %%
STATIC_LEGEND_HANDLES = [
    Line2D(
        [0],
        [0],
        marker="o",
        color="w",
        markerfacecolor=COLORS["blue"],
        markersize=10,
        label="Companies",
    ),
    Line2D(
        [0],
        [0],
        marker="o",
        color="w",
        markerfacecolor=COLORS["amber"],
        markersize=14,
        label="Shared suppliers",
    ),
]

# %% [markdown]
# ### Static Footer
#
# Add the book attribution separately so the annotation helper remains compact.


# %%
def add_static_footer(ax: plt.Axes) -> None:
    """Add the ML4T attribution footer."""
    ax.text(
        0.99,
        0.01,
        "ML4T 3rd Edition",
        transform=ax.transAxes,
        fontsize=8,
        ha="right",
        va="bottom",
        color=COLORS["slate"],
        alpha=0.7,
    )


# %% [markdown]
# ### Static Annotation Helper
#
# Apply the callout, title, legend, and footer after plotting the network.


# %%
def finalize_static_figure(
    ax: plt.Axes,
    G: nx.Graph,
    pos: dict[str, tuple[float, float]],
    sorted_suppliers: list[str],
) -> None:
    """Add the supplier callout, title, legend, and footer."""
    add_static_supplier_callout(ax, G, pos, sorted_suppliers)
    add_message_title(
        ax,
        "Suppliers named by several companies, and who names them",
        subtitle=(
            f"{EXTRACTOR_NAME} extraction of S&P 100 10-K excerpts; inner ring sized "
            "by the number of companies naming the supplier"
        ),
    )
    ax.legend(handles=STATIC_LEGEND_HANDLES, loc="upper left", frameon=True)
    add_static_footer(ax)


# %% [markdown]
# ### Static Figure Function
#
# Render the network using the precomputed layout and then add the annotation
# and legend.


# %%
def create_static_figure(G: nx.Graph) -> plt.Figure:
    """Create publication-ready supply chain network figure."""
    fig, ax = plt.subplots(figsize=FIGSIZE["single_tall"], constrained_layout=True)
    _, companies, pos, sorted_suppliers = build_static_layout(G)
    nx.draw_networkx_edges(G, pos, ax=ax, edge_color=COLORS["silver_muted"], width=0.8, alpha=0.6)
    nx.draw_networkx_nodes(
        G,
        pos,
        nodelist=companies,
        ax=ax,
        node_color=COLORS["blue"],
        node_size=80,
        alpha=0.8,
    )
    nx.draw_networkx_nodes(
        G,
        pos,
        nodelist=sorted_suppliers,
        ax=ax,
        node_color=COLORS["amber"],
        node_size=[100 + G.degree(n) * 20 for n in sorted_suppliers],
        alpha=0.9,
    )
    top_supplier_labels = {n: n for n in sorted_suppliers[:6]}
    nx.draw_networkx_labels(
        G, pos, labels=top_supplier_labels, ax=ax, font_size=7, font_weight="bold"
    )
    finalize_static_figure(ax, G, pos, sorted_suppliers)
    ax.axis("off")
    return fig


# %%
# Render the static figure inline. The publication-quality PNG/PDF
# (figure_23_3_supply_chain_network.*) is generated by the book-repo figure
# scripts, not saved here. Notebooks display figures without writing publication files.
OUTPUT_ROOT = Path(os.getenv("ML4T_OUTPUT_DIR", get_chapter_dir(23) / "output"))
VIZ_DIR = OUTPUT_ROOT / "ch23" / "supply_chain_visualizations"
VIZ_DIR.mkdir(parents=True, exist_ok=True)

static_fig = create_static_figure(G)
top_shared_name, top_shared_companies = max(shared_suppliers.items(), key=lambda kv: len(kv[1]))
show_with_alt(
    static_fig,
    f"A two-ring network diagram. The inner ring holds the {len(network_suppliers)} "
    f"suppliers named by more than one company, the largest being {top_shared_name} at "
    f"{len(top_shared_companies)} companies. The outer ring holds the "
    f"{len(network_companies)} companies that name at least one of them, each joined "
    "by a line to the suppliers it names.",
)

# %% [markdown]
# Larger inner-ring nodes are named by more companies in the extracted candidate
# graph. This shared-neighbor structure is directly queryable in a graph
# representation. The figure omits every supplier named by a single company, which
# is almost all of them, so it shows where the graph has shared structure rather
# than what the graph mostly contains. A disruption scenario still needs event
# evidence and a model that links exposure to portfolio outcomes.

# %% [markdown]
# ### 12.2 Interactive Notebook Visualization
#
# pyvis network for interactive exploration (hover, zoom, drag). Nodes display
# supplier risk levels (HIGH/MEDIUM/LOW based on dependency count) on hover.

# %%
from pyvis.network import Network

# %% [markdown]
# ### Interactive Graph Builder
#
# Creates a pyvis force-directed graph with tooltips showing concentration risk
# levels and supplier details for each node.

# %% [markdown]
# ### Interactive Node Labels
#
# Use a shared risk classifier so supplier tooltips and the D3 export apply the
# same high/medium/low thresholds.


# %% [markdown]
# Thresholds on the number of companies naming one supplier. They are cut points for
# a display, not a calibrated risk scale: nothing here estimates a disruption
# probability, and the counts they read come from what filings happen to name.

# %%
RISK_HIGH_COMPANIES = 10
RISK_MEDIUM_COMPANIES = 5


def supplier_risk_level(count: int) -> str:
    """Bucket a supplier by how many companies name it, for notebook displays."""
    if count >= RISK_HIGH_COMPANIES:
        return "HIGH"
    if count >= RISK_MEDIUM_COMPANIES:
        return "MEDIUM"
    return "LOW"


# %% [markdown]
# ### Interactive Graph Builder
#
# Create the pyvis network with supplier risk tooltips and compact company
# labels for notebook exploration.

# %% [markdown]
# ### Interactive Node Helper
#
# Add one node at a time so the main pyvis builder focuses on orchestration.


# %%
def add_interactive_node(net: Network, G: nx.Graph, node: str, data: dict[str, object]) -> None:
    """Add a supplier or company node with its tooltip."""
    if data.get("node_type") == "supplier":
        count = G.degree(node)
        title = (
            f"<b>{node}</b><br>Supplies {count} companies"
            f"<br><i>Concentration risk: {supplier_risk_level(count)}</i>"
        )
        net.add_node(node, label=node, color=COLORS["amber"], size=15 + count * 3, title=title)
        return
    neighbors = list(G.neighbors(node))
    supplier_list = ", ".join(neighbors[:5])
    title = f"<b>{node}</b><br>Key suppliers: {supplier_list}"
    net.add_node(node, label=node[:12], color=COLORS["blue"], size=20, title=title)


# %% [markdown]
# #

Shown in full with attribution under the source's licence. Licence: MIT

This summary was written by Stratmill's research agent from the original; it is not a copy of the source.