From 0f0800f10953cfdd41df1617e7272a1dc1844d66 Mon Sep 17 00:00:00 2001 From: KaifAhmad1 Date: Thu, 26 Mar 2026 14:20:38 +0530 Subject: [PATCH] =?UTF-8?q?feat(#402):=20Temporal=20GraphRAG=20Integration?= =?UTF-8?q?=20=E2=80=94=20TemporalGraphRetriever=20&=20TemporalQueryRewrit?= =?UTF-8?q?er?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add TemporalGraphRetriever to context_retriever.py (no new file per project convention) - Drop-in wrapper for ContextRetriever; filters related_entities/related_relationships via reconstruct_at_time(); at_time=None is a true passthrough - Returns new RetrievedContext objects (no in-place mutation) - Graceful ImportError if temporal modules unavailable - Add at_time + header_template to ContextRetriever._generate_reasoned_response() and query_with_reasoning() - Temporal header prepended to LLM context block only when at_time is set - Naive datetimes normalised to UTC before formatting - Header built with str.replace (not .format) to prevent format-string injection - Add TemporalQueryRewriter + TemporalQueryResult to semantica/kg/ - Regex-only (default) and LLM-assisted extraction modes - Resolves temporal phrases via TemporalNormalizer (deterministic, zero LLM) - Word-boundary guards on intent keywords; year fallback for noun-phrase dates - Never calls reconstruct_at_time — extraction only - Export TemporalGraphRetriever from semantica.context - Export TemporalQueryRewriter, TemporalQueryResult from semantica.kg - Add 99 tests across two new test files - tests/context/test_temporal_retriever.py (56 tests) - tests/kg/test_temporal_query_rewriter.py (43 tests) Co-Authored-By: Claude Sonnet 4.6 --- semantica/context/__init__.py | 3 +- semantica/context/context_retriever.py | 215 ++++++++++- semantica/kg/__init__.py | 3 + semantica/kg/temporal_query_rewriter.py | 457 +++++++++++++++++++++++ tests/context/test_temporal_retriever.py | 428 +++++++++++++++++++++ tests/kg/test_temporal_query_rewriter.py | 418 +++++++++++++++++++++ 6 files changed, 1517 insertions(+), 7 deletions(-) create mode 100644 semantica/kg/temporal_query_rewriter.py create mode 100644 tests/context/test_temporal_retriever.py create mode 100644 tests/kg/test_temporal_query_rewriter.py diff --git a/semantica/context/__init__.py b/semantica/context/__init__.py index ff50605d..c522e4f7 100644 --- a/semantica/context/__init__.py +++ b/semantica/context/__init__.py @@ -108,7 +108,7 @@ Production Examples: from .agent_context import AgentContext from .agent_memory import AgentMemory, MemoryItem from .context_graph import ContextEdge, ContextGraph, ContextNode -from .context_retriever import ContextRetriever, RetrievedContext +from .context_retriever import ContextRetriever, RetrievedContext, TemporalGraphRetriever from .decision_context import DecisionContext from .entity_linker import EntityLink, EntityLinker, LinkedEntity @@ -144,6 +144,7 @@ __all__ = [ "MemoryItem", "ContextRetriever", "RetrievedContext", + "TemporalGraphRetriever", # Decision tracking models "Decision", "DecisionContextModel", diff --git a/semantica/context/context_retriever.py b/semantica/context/context_retriever.py index 52752f75..b7adec96 100644 --- a/semantica/context/context_retriever.py +++ b/semantica/context/context_retriever.py @@ -76,7 +76,7 @@ Author: Semantica Contributors License: MIT """ -from datetime import datetime +from datetime import datetime, timezone from dataclasses import dataclass, field from typing import Any, Dict, List, Optional, Union @@ -88,6 +88,14 @@ from ..kg.path_finder import PathFinder from ..kg.centrality_calculator import CentralityCalculator from ..kg.community_detector import CommunityDetector from ..kg.similarity_calculator import SimilarityCalculator +try: + from ..kg.temporal_query import TemporalGraphQuery as _TemporalGraphQuery + from ..kg.temporal_model import parse_temporal_value as _parse_temporal_value + _TEMPORAL_AVAILABLE = True +except Exception: # pragma: no cover + _TemporalGraphQuery = None # type: ignore[assignment,misc] + _parse_temporal_value = None # type: ignore[assignment] + _TEMPORAL_AVAILABLE = False @dataclass @@ -1381,7 +1389,9 @@ class ContextRetriever: query: str, retrieved_context: List[RetrievedContext], reasoning_paths: List[Dict[str, Any]], - llm_provider: Any + llm_provider: Any, + at_time: Optional[Any] = None, + header_template: str = "[Graph context valid as of: {at_time} UTC | Source: {source}]", ) -> str: """ Generate natural language response using LLM with retrieved context and reasoning paths. @@ -1391,10 +1401,36 @@ class ContextRetriever: retrieved_context: Retrieved context items reasoning_paths: Multi-hop reasoning paths llm_provider: LLM provider instance (from semantica.llms) + at_time: Optional point-in-time for the graph snapshot. When set, a + structured temporal header is prepended to the context block so + the LLM knows the validity window of the facts it is reasoning over. + header_template: Format string for the temporal header. Placeholders: + ``{at_time}`` (ISO timestamp) and ``{source}`` (snapshot label). Returns: Generated natural language response """ + # Build temporal header when at_time is explicitly provided + temporal_header = "" + if at_time is not None: + if isinstance(at_time, datetime): + # Normalise to UTC so the header always says "UTC" truthfully + if at_time.tzinfo is None: + at_time = at_time.replace(tzinfo=timezone.utc) + else: + at_time = at_time.astimezone(timezone.utc) + at_time_str = at_time.isoformat() + else: + at_time_str = str(at_time) + # Use plain str.replace instead of .format() so that unexpected + # braces in header_template or in at_time_str cannot be + # interpreted as additional format placeholders (injection guard). + temporal_header = ( + header_template + .replace("{at_time}", at_time_str) + .replace("{source}", "KnowledgeGraph snapshot") + ) + "\n\n" + # Format retrieved context context_text = "\n\n".join([ f"Context {i+1} (Score: {ctx.score:.2f}):\n{ctx.content}" @@ -1426,7 +1462,7 @@ class ContextRetriever: User Question: {query} -Retrieved Context: +{temporal_header}Retrieved Context: {context_text} {reasoning_text} @@ -1454,7 +1490,9 @@ Answer:""" llm_provider: Any, max_results: int = 10, max_hops: int = 2, - **kwargs + at_time: Optional[Any] = None, + header_template: str = "[Graph context valid as of: {at_time} UTC | Source: {source}]", + **kwargs, ) -> Dict[str, Any]: """ Query with multi-hop reasoning and LLM-based response generation. @@ -1467,7 +1505,16 @@ Answer:""" llm_provider: LLM provider instance (from semantica.llms) max_results: Maximum context results to retrieve (default: 10) max_hops: Maximum graph traversal hops (default: 2) - **kwargs: Additional retrieval options + at_time: Optional point-in-time for the graph snapshot. Accepts a + :class:`~datetime.datetime` (naive datetimes are assumed UTC) + or an ISO 8601 string. When set, a structured temporal header + is prepended to the LLM context block so the model knows the + validity window of the facts it reasons over. + header_template: Template string for the temporal header. Only + ``{at_time}`` and ``{source}`` placeholders are substituted; + any other braces are left as-is. Defaults to + ``"[Graph context valid as of: {at_time} UTC | Source: {source}]"``. + **kwargs: Additional retrieval options passed to ``retrieve()`` Returns: Dictionary with: @@ -1541,7 +1588,9 @@ Answer:""" query, retrieved_context, reasoning_paths, - llm_provider + llm_provider, + at_time=at_time, + header_template=header_template, ) # Step 5: Format reasoning path as string @@ -2596,3 +2645,157 @@ Answer:""" return DecisionQuery(self.knowledge_graph) except Exception: return None + + +class TemporalGraphRetriever: + """ + Drop-in temporal wrapper for any ContextRetriever. + + Calls ``base_retriever.retrieve(query)``, then filters ``related_entities`` + and ``related_relationships`` in each :class:`RetrievedContext` item to + those active at *at_time* using + :meth:`~semantica.kg.temporal_query.TemporalGraphQuery.reconstruct_at_time`. + + When no *at_time* is set (neither in the constructor nor at the call site) + the base retriever result is returned unchanged — no copies, no filtering. + + Example:: + + from semantica.context import ContextRetriever, TemporalGraphRetriever + + base = ContextRetriever(knowledge_graph=kg) + retriever = TemporalGraphRetriever(base, at_time="2023-06-01") + results = retriever.retrieve("which suppliers were certified?") + # related_entities / related_relationships in each result are valid as + # of 2023-06-01; dangling edges are removed automatically. + """ + + _DEFAULT_HEADER = "[Graph context valid as of: {at_time} UTC | Source: {source}]" + + def __init__( + self, + base_retriever: "ContextRetriever", + at_time: Optional[Any] = None, + header_template: str = _DEFAULT_HEADER, + ): + """ + Args: + base_retriever: Any :class:`ContextRetriever` instance to wrap. + at_time: Default point-in-time for temporal filtering. Accepts a + :class:`~datetime.datetime` or any string parseable by + :func:`~semantica.kg.temporal_model.parse_temporal_value` + (e.g. ``"2023-06-01"``). ``None`` disables filtering. + header_template: Format string used to build the temporal context + header injected into LLM prompts by + :meth:`ContextRetriever.query_with_reasoning`. Placeholders: + ``{at_time}`` and ``{source}``. + + Example: + >>> from semantica.context import ContextRetriever, TemporalGraphRetriever + + >>> # Passthrough — no temporal filtering applied + >>> base = ContextRetriever(knowledge_graph=kg) + >>> retriever = TemporalGraphRetriever(base) + >>> results = retriever.retrieve("active suppliers") # identical to base.retrieve() + + >>> # Point-in-time snapshot — ISO string shorthand + >>> retriever = TemporalGraphRetriever(base, at_time="2023-06-01") + >>> results = retriever.retrieve("certified suppliers") + >>> # related_entities/related_relationships are valid as of 2023-06-01 + + >>> # Custom prompt header for LLM context + >>> retriever = TemporalGraphRetriever( + ... base, + ... at_time="2023-06-01", + ... header_template="[Snapshot: {at_time} | {source}]", + ... ) + + >>> # Combine with TemporalQueryRewriter for end-to-end temporal RAG + >>> from semantica.kg import TemporalQueryRewriter + >>> rw = TemporalQueryRewriter() + >>> parsed = rw.rewrite("which suppliers were certified before 2022?") + >>> retriever = TemporalGraphRetriever(base, at_time=parsed.at_time) + >>> results = retriever.retrieve(parsed.rewritten_query) + """ + if not _TEMPORAL_AVAILABLE: + raise ImportError( + "TemporalGraphRetriever requires the semantica.kg temporal modules " + "(temporal_query, temporal_model). Ensure they are installed and " + "importable without errors." + ) + self.base_retriever = base_retriever + self.at_time = at_time + self.header_template = header_template + self._tgq = _TemporalGraphQuery() + + def retrieve( + self, + query: str, + at_time: Optional[Any] = None, + **kwargs, + ) -> List[RetrievedContext]: + """ + Retrieve context and apply point-in-time temporal filtering. + + Args: + query: Search query forwarded to the base retriever. + at_time: Override the instance-level ``at_time`` for this call. + ``None`` falls back to the constructor value; if both are + ``None`` the base result is returned unchanged. + **kwargs: Forwarded to ``base_retriever.retrieve()``. + + Returns: + List of :class:`RetrievedContext` items whose + ``related_entities`` and ``related_relationships`` are restricted + to facts valid at *at_time*. Dangling relationships (whose + source or target entity was filtered out) are removed. + + Example: + >>> from datetime import datetime, timezone + >>> from semantica.context import ContextRetriever, TemporalGraphRetriever + + >>> base = ContextRetriever(knowledge_graph=kg) + >>> retriever = TemporalGraphRetriever(base, at_time="2023-06-01") + + >>> # Use constructor at_time + >>> results = retriever.retrieve("drug interactions") + >>> for r in results: + ... print(len(r.related_entities), "entities valid as of 2023-06-01") + + >>> # Override at_time per call (e.g. from TemporalQueryRewriter output) + >>> results_q1 = retriever.retrieve( + ... "capital requirements", + ... at_time=datetime(2023, 3, 31, tzinfo=timezone.utc), + ... ) + + >>> # Passthrough: no at_time → base result unchanged + >>> retriever_plain = TemporalGraphRetriever(base) + >>> plain = retriever_plain.retrieve("suppliers") # no temporal filtering + """ + import dataclasses + + effective_at_time = at_time if at_time is not None else self.at_time + results = self.base_retriever.retrieve(query, **kwargs) + if effective_at_time is None: + return results + parsed = ( + effective_at_time + if isinstance(effective_at_time, datetime) + else _parse_temporal_value(effective_at_time) + ) + filtered_results = [] + for ctx in results: + subgraph = { + "entities": ctx.related_entities, + "relationships": ctx.related_relationships, + } + filtered = self._tgq.reconstruct_at_time(subgraph, parsed) + # Return a new RetrievedContext rather than mutating the original + # so callers that hold a reference to the base retriever's results + # are not surprised by side-effects. + filtered_results.append(dataclasses.replace( + ctx, + related_entities=filtered["entities"], + related_relationships=filtered["relationships"], + )) + return filtered_results diff --git a/semantica/kg/__init__.py b/semantica/kg/__init__.py index e2f5a5f2..c31912ef 100644 --- a/semantica/kg/__init__.py +++ b/semantica/kg/__init__.py @@ -128,6 +128,7 @@ from .temporal_query import ( ) from .temporal_model import BiTemporalFact, TemporalBound from .temporal_normalizer import TemporalNormalizer +from .temporal_query_rewriter import TemporalQueryRewriter, TemporalQueryResult __all__ = [ # Core Classes @@ -142,6 +143,8 @@ __all__ = [ "TemporalBound", "BiTemporalFact", "TemporalNormalizer", + "TemporalQueryRewriter", + "TemporalQueryResult", "AlgorithmTrackerWithProvenance", "ProvenanceTracker", # Enhanced Graph Algorithms diff --git a/semantica/kg/temporal_query_rewriter.py b/semantica/kg/temporal_query_rewriter.py new file mode 100644 index 00000000..8b37f5b7 --- /dev/null +++ b/semantica/kg/temporal_query_rewriter.py @@ -0,0 +1,457 @@ +""" +Temporal Query Rewriter + +Extracts temporal references from natural-language queries so that downstream +components (e.g. :class:`~semantica.context.context_retriever.TemporalGraphRetriever`) +can perform **deterministic** temporal filtering instead of asking the LLM to +do it. + +Separation of concerns +----------------------- +* **This module** — extracts ``at_time``, ``start_time``, ``end_time``, and + ``temporal_intent`` from query text and produces a cleaned ``rewritten_query`` + with the temporal phrase stripped out. +* **TemporalGraphRetriever** — uses the extracted parameters to call + ``reconstruct_at_time()`` on the retrieved subgraph. + +This module never calls ``reconstruct_at_time()``. + +Usage:: + + from semantica.kg import TemporalQueryRewriter + + rewriter = TemporalQueryRewriter() # no LLM — regex only + result = rewriter.rewrite("which suppliers were certified before the 2021 merger?") + # result.temporal_intent == "before" + # result.at_time == datetime(2021, 1, 1, tzinfo=utc) (start of year) + # result.rewritten_query == "which suppliers were certified?" + + # With an LLM for higher accuracy on free-form phrasing: + from semantica.llms import Groq + llm = Groq(model="llama-3.1-8b-instant") + rewriter = TemporalQueryRewriter(llm_provider=llm) + result = rewriter.rewrite("what interactions were known in Q2 2022?") + # result.temporal_intent == "during" + # result.start_time / result.end_time == Q2 2022 bounds +""" + +from __future__ import annotations + +import json +import logging +import re +from dataclasses import dataclass, field +from datetime import datetime, timezone +from typing import Any, Dict, Optional, Tuple + +from .temporal_normalizer import TemporalNormalizer + +logger = logging.getLogger(__name__) + + +# --------------------------------------------------------------------------- +# Result dataclass +# --------------------------------------------------------------------------- + +@dataclass +class TemporalQueryResult: + """ + Output of :meth:`TemporalQueryRewriter.rewrite`. + + Attributes: + rewritten_query: Original query with the temporal phrase removed and + whitespace normalised. Identical to the input query when no + temporal phrase is found. + at_time: Point-in-time bound for intents ``"at"``, ``"during"``, + ``"before"``, and ``"after"``. For ``"before"`` / ``"after"`` + this is the bounding moment. ``None`` for ``"between"`` (use + ``start_time`` / ``end_time``) and when no temporal phrase is + found. + start_time: Lower bound for ``"between"`` queries. ``None`` otherwise. + end_time: Upper bound for ``"between"`` queries. ``None`` otherwise. + temporal_intent: One of ``"before"``, ``"after"``, ``"at"``, + ``"during"``, ``"between"``, or ``None`` when no temporal phrase + was detected. + confidence: Extraction confidence in ``[0.0, 1.0]``. Regex-only + extractions return ``0.85``; LLM-backed extractions propagate + the LLM's self-reported confidence or default to ``0.75``. + """ + + rewritten_query: str + at_time: Optional[datetime] = None + start_time: Optional[datetime] = None + end_time: Optional[datetime] = None + temporal_intent: Optional[str] = None + confidence: float = 0.0 + + def has_temporal_context(self) -> bool: + """Return ``True`` when any temporal parameter was extracted. + + Example: + >>> from semantica.kg import TemporalQueryRewriter + >>> rw = TemporalQueryRewriter() + >>> result = rw.rewrite("suppliers certified before 2021") + >>> result.has_temporal_context() + True + >>> rw.rewrite("list all suppliers").has_temporal_context() + False + """ + return self.temporal_intent is not None + + +# --------------------------------------------------------------------------- +# Regex-based fallback extractor +# --------------------------------------------------------------------------- + +# Maps leading keyword → temporal_intent +_INTENT_PREFIXES: list[Tuple[str, str]] = [ + # "between … and …" handled separately + (r"\bbetween\b", "between"), + (r"\bprior\s+to\b", "before"), + (r"\bbefore\b", "before"), + (r"\buntil\b", "before"), + (r"\bup\s+to\b", "before"), + (r"\bafter\b", "after"), + (r"\bsince\b", "after"), + (r"\bfollowing\b", "after"), + (r"\bduring\b", "during"), + (r"\bin\b", "during"), + (r"\bwithin\b", "during"), + (r"\bas\s+of\b", "at"), + (r"\bat\b", "at"), + (r"\bon\b", "at"), +] + +# Captures a temporal phrase after an intent keyword +_TEMPORAL_PHRASE_RE = re.compile( + r""" + \b + (?P + between | prior\s+to | before | until | up\s+to | + after | since | following | + during | in | within | + as\s+of | at | on + ) + \b + \s+ + (?P + # "between X and Y" + (?: + (?:the\s+)? + [\w\s,\-/]+ (?:\s+and\s+[\w\s,\-/]+)? + ) + | + # simple phrase: "Q1 2022", "2021 merger", "June 2020" + [\w][\w\s,\-/]*? + ) + (?=\s*[?.,;!]|\s*$|\s+(?:and\b|\bor\b|\bthat\b|\bwhere\b|\bwho\b|\bwhich\b|\bwhen\b)) + """, + re.IGNORECASE | re.VERBOSE, +) + +# "between X and Y" extractor +_BETWEEN_RE = re.compile( + r"\bbetween\s+(?P[\w][\w\s\-,/]*?)\s+and\s+(?P[\w][\w\s\-,/]*?)" + r"(?=\s*[?.,;!]|\s*$|\s+(?:and\b|\bor\b|\bthat\b|\bwhere\b|\bwho\b|\bwhich\b|\bwhen\b))", + re.IGNORECASE, +) + + +def _map_intent_keyword(kw: str) -> str: + kw_lower = kw.lower().strip() + if kw_lower in ("before", "until", "up to", "prior to"): + return "before" + if kw_lower in ("after", "since", "following"): + return "after" + if kw_lower in ("during", "in", "within"): + return "during" + if kw_lower in ("as of", "at", "on"): + return "at" + return "at" + + +def _strip_phrase(query: str, matched_text: str) -> str: + """Remove *matched_text* from *query* and tidy up residual whitespace.""" + cleaned = query.replace(matched_text, " ") + cleaned = re.sub(r"\s{2,}", " ", cleaned).strip() + cleaned = re.sub(r"\s+([?.,;!])", r"\1", cleaned) + return cleaned + + +# --------------------------------------------------------------------------- +# Main class +# --------------------------------------------------------------------------- + +class TemporalQueryRewriter: + """ + Extract temporal intent from natural-language queries. + + Two operating modes: + + **Regex-only** (default, no LLM required) + Handles common structured phrasings — ``"before 2021"``, + ``"between Q1 and Q3 2022"``, ``"as of 2023-06-01"``, etc. + All datetime resolution is delegated to + :class:`~semantica.kg.temporal_normalizer.TemporalNormalizer` + (deterministic, zero LLM calls). + + **LLM-assisted** (``llm_provider`` kwarg) + Uses a small LLM call to extract the temporal phrase and intent for + free-form phrasing. Datetime resolution is still performed by + :class:`~semantica.kg.temporal_normalizer.TemporalNormalizer` so the + result is always deterministic given the same phrase text. + + .. warning:: + This class never calls ``reconstruct_at_time()``. It is a pure + parameter-extraction step; the actual temporal filtering is always + performed by + :class:`~semantica.context.context_retriever.TemporalGraphRetriever`. + """ + + _LLM_EXTRACTION_PROMPT = ( + "Extract the temporal reference from this query and return ONLY valid JSON " + "with these exact keys:\n" + ' "temporal_phrase": the verbatim temporal expression (or null),\n' + ' "temporal_intent": one of "before", "after", "at", "during", "between", or null,\n' + ' "confidence": a float between 0 and 1.\n\n' + "Query: {query}\n\nJSON:" + ) + + def __init__( + self, + llm_provider: Optional[Any] = None, + reference_date: Optional[datetime] = None, + ): + """ + Args: + llm_provider: Optional LLM provider (any object with a + ``generate(prompt: str) -> str`` method). When ``None``, + regex-only extraction is used. + reference_date: Reference date for relative phrases like + ``"last year"``. Defaults to ``datetime.now(utc)`` at + construction time. + + Example: + >>> from semantica.kg import TemporalQueryRewriter + >>> # Regex-only (no LLM dependency) + >>> rw = TemporalQueryRewriter() + + >>> # LLM-assisted — handles free-form phrasings like "the 2021 merger" + >>> from semantica.llms import Groq + >>> llm = Groq(model="llama-3.1-8b-instant") + >>> rw_llm = TemporalQueryRewriter(llm_provider=llm) + """ + self._llm = llm_provider + self._normalizer = TemporalNormalizer( + reference_date=reference_date or datetime.now(timezone.utc) + ) + + # ------------------------------------------------------------------ + # Public API + # ------------------------------------------------------------------ + + def rewrite( + self, + query: str, + context: Optional[Dict[str, Any]] = None, + ) -> TemporalQueryResult: + """ + Extract temporal context from *query* and return a cleaned query. + + Args: + query: The raw user query. + context: Optional dict of additional context (e.g. domain hint). + Currently unused; reserved for future domain-specific rules. + + Returns: + :class:`TemporalQueryResult` with extracted parameters and the + cleaned ``rewritten_query``. + + Example: + >>> from semantica.kg import TemporalQueryRewriter + >>> rw = TemporalQueryRewriter() + + >>> # "before" intent — single bound + >>> r = rw.rewrite("which suppliers were certified before 2021?") + >>> r.temporal_intent + 'before' + >>> r.at_time.year + 2021 + >>> r.rewritten_query + 'which suppliers were certified?' + + >>> # "between" intent — range + >>> r = rw.rewrite("revenue between Q1 2022 and Q3 2022") + >>> r.temporal_intent + 'between' + >>> r.start_time is not None and r.end_time is not None + True + + >>> # No temporal phrase — passthrough + >>> r = rw.rewrite("list all active suppliers") + >>> r.temporal_intent is None + True + >>> r.rewritten_query + 'list all active suppliers' + """ + if self._llm is not None: + result = self._llm_rewrite(query) + if result is not None: + return result + return self._regex_rewrite(query) + + # ------------------------------------------------------------------ + # Internal: LLM path + # ------------------------------------------------------------------ + + def _llm_rewrite(self, query: str) -> Optional[TemporalQueryResult]: + """ + Ask the LLM to extract the temporal phrase + intent. + + Falls back to ``None`` (triggering the regex path) if the LLM call + fails or produces unparseable output. + """ + prompt = self._LLM_EXTRACTION_PROMPT.format(query=query) + try: + raw = self._llm.generate(prompt) + data = self._parse_json(raw) + except Exception as exc: + logger.debug("LLM extraction failed (%s); falling back to regex.", exc) + return None + + phrase: Optional[str] = data.get("temporal_phrase") + intent: Optional[str] = data.get("temporal_intent") + confidence: float = float(data.get("confidence", 0.75)) + + if not phrase or not intent: + return TemporalQueryResult( + rewritten_query=query, + temporal_intent=None, + confidence=confidence, + ) + + at_time, start_time, end_time = self._resolve_phrase(phrase, intent) + rewritten = _strip_phrase(query, phrase) + + return TemporalQueryResult( + rewritten_query=rewritten, + at_time=at_time, + start_time=start_time, + end_time=end_time, + temporal_intent=intent, + confidence=confidence, + ) + + @staticmethod + def _parse_json(text: str) -> Dict[str, Any]: + """Extract JSON object from raw LLM output (tolerates trailing prose).""" + # Find first '{' … '}' block + start = text.find("{") + end = text.rfind("}") + 1 + if start == -1 or end == 0: + raise ValueError("No JSON object found in LLM output") + return json.loads(text[start:end]) + + # ------------------------------------------------------------------ + # Internal: regex path + # ------------------------------------------------------------------ + + def _regex_rewrite(self, query: str) -> TemporalQueryResult: + """Regex-based extraction — no LLM calls.""" + # 1. Try "between X and Y" first (most specific) + between_m = _BETWEEN_RE.search(query) + if between_m: + start_phrase = between_m.group("start").strip() + end_phrase = between_m.group("end").strip() + start_time = self._normalize_single(start_phrase, prefer="start") + end_time = self._normalize_single(end_phrase, prefer="end") + if start_time is not None or end_time is not None: + rewritten = _strip_phrase(query, between_m.group(0)) + return TemporalQueryResult( + rewritten_query=rewritten, + start_time=start_time, + end_time=end_time, + temporal_intent="between", + confidence=0.85, + ) + + # 2. Generic intent + phrase + m = _TEMPORAL_PHRASE_RE.search(query) + if m: + intent_kw = m.group("intent_kw") + phrase = m.group("phrase").strip() + intent = _map_intent_keyword(intent_kw) + + at_time, start_time, end_time = self._resolve_phrase(phrase, intent) + if at_time is not None or start_time is not None or end_time is not None: + rewritten = _strip_phrase(query, m.group(0)) + return TemporalQueryResult( + rewritten_query=rewritten, + at_time=at_time, + start_time=start_time, + end_time=end_time, + temporal_intent=intent, + confidence=0.85, + ) + + # 3. No temporal phrase found + return TemporalQueryResult( + rewritten_query=query, + temporal_intent=None, + confidence=1.0, + ) + + # ------------------------------------------------------------------ + # Internal: phrase → datetime resolution + # ------------------------------------------------------------------ + + # Matches a 4-digit year as fallback when the full phrase can't be normalized + _YEAR_IN_PHRASE = re.compile(r"\b(\d{4})\b") + + def _resolve_phrase( + self, + phrase: str, + intent: str, + ) -> Tuple[Optional[datetime], Optional[datetime], Optional[datetime]]: + """ + Resolve *phrase* to datetime bounds via :class:`TemporalNormalizer`. + + Returns ``(at_time, start_time, end_time)`` according to *intent*: + + * ``"before"`` / ``"after"`` / ``"at"`` / ``"during"`` → only ``at_time`` + is populated (``start_time`` / ``end_time`` are ``None``). + * ``"between"`` → only ``start_time`` / ``end_time`` are populated. + + When the full phrase cannot be normalized (e.g. ``"the 2021 merger"``) + the method extracts the first 4-digit year as a fallback and retries. + """ + try: + norm_start, norm_end = self._normalizer.normalize(phrase) + except Exception: + # Fallback: extract first 4-digit year from phrase and retry + year_m = self._YEAR_IN_PHRASE.search(phrase) + if year_m: + try: + norm_start, norm_end = self._normalizer.normalize(year_m.group(1)) + except Exception: + return None, None, None + else: + return None, None, None + + if intent == "between": + return None, norm_start, norm_end + + # For all point-in-time intents use the start of the resolved interval + # as the bounding datetime ("before 2021" → before 2021-01-01). + at_time = norm_start + return at_time, None, None + + def _normalize_single( + self, phrase: str, prefer: str = "start" + ) -> Optional[datetime]: + """Normalize *phrase* to a single datetime (start or end of interval).""" + try: + start, end = self._normalizer.normalize(phrase) + return start if prefer == "start" else end + except Exception: + return None diff --git a/tests/context/test_temporal_retriever.py b/tests/context/test_temporal_retriever.py new file mode 100644 index 00000000..35ecddea --- /dev/null +++ b/tests/context/test_temporal_retriever.py @@ -0,0 +1,428 @@ +""" +Tests for TemporalGraphRetriever and temporal context header in LLM prompts. +""" + +from datetime import datetime, timezone +from unittest.mock import MagicMock, call, patch + +import pytest + +from semantica.context import ContextRetriever, RetrievedContext, TemporalGraphRetriever + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _utc(year, month=1, day=1): + return datetime(year, month, day, tzinfo=timezone.utc) + + +def _make_entity(eid, valid_from=None, valid_until=None, **extra): + e = {"id": eid, "name": eid} + if valid_from: + e["valid_from"] = valid_from.isoformat() + if valid_until: + e["valid_until"] = valid_until.isoformat() + e.update(extra) + return e + + +def _make_rel(source, target, rel_type="RELATED_TO", valid_from=None, valid_until=None): + r = {"source": source, "target": target, "type": rel_type} + if valid_from: + r["valid_from"] = valid_from.isoformat() + if valid_until: + r["valid_until"] = valid_until.isoformat() + return r + + +def _make_result(entities, relationships, content="test content", score=0.9): + return RetrievedContext( + content=content, + score=score, + source="graph", + related_entities=entities, + related_relationships=relationships, + ) + + +def _base_mock(return_value=None): + base = MagicMock(spec=ContextRetriever) + base.retrieve.return_value = return_value or [] + return base + + +# --------------------------------------------------------------------------- +# TemporalGraphRetriever — construction +# --------------------------------------------------------------------------- + +class TestTemporalGraphRetrieverInit: + + def test_default_at_time_is_none(self): + tr = TemporalGraphRetriever(_base_mock()) + assert tr.at_time is None + + def test_stores_base_retriever(self): + base = _base_mock() + tr = TemporalGraphRetriever(base) + assert tr.base_retriever is base + + def test_custom_header_template_stored(self): + tr = TemporalGraphRetriever(_base_mock(), header_template="[{at_time}]") + assert tr.header_template == "[{at_time}]" + + def test_default_header_template_contains_placeholder(self): + tr = TemporalGraphRetriever(_base_mock()) + assert "{at_time}" in tr.header_template + assert "{source}" in tr.header_template + + +# --------------------------------------------------------------------------- +# TemporalGraphRetriever — passthrough (no at_time) +# --------------------------------------------------------------------------- + +class TestTemporalGraphRetrieverPassthrough: + + def setup_method(self): + self.entity = _make_entity("e1", _utc(2020, 1, 1), _utc(2025, 1, 1)) + self.base = _base_mock([_make_result([self.entity], [])]) + + def test_no_at_time_returns_base_result_unchanged(self): + tr = TemporalGraphRetriever(self.base) + results = tr.retrieve("some query") + self.base.retrieve.assert_called_once_with("some query") + assert results[0].related_entities[0]["id"] == "e1" + + def test_none_at_time_on_call_uses_constructor_none(self): + tr = TemporalGraphRetriever(self.base, at_time=None) + results = tr.retrieve("query", at_time=None) + assert results[0].related_entities[0]["id"] == "e1" + + def test_empty_base_result_passthrough(self): + base = _base_mock([]) + tr = TemporalGraphRetriever(base) + assert tr.retrieve("q") == [] + + def test_kwargs_forwarded_to_base_retriever(self): + tr = TemporalGraphRetriever(self.base) + tr.retrieve("q", max_results=3, min_relevance_score=0.5) + self.base.retrieve.assert_called_once_with("q", max_results=3, min_relevance_score=0.5) + + def test_base_retriever_called_exactly_once(self): + tr = TemporalGraphRetriever(self.base) + tr.retrieve("q") + assert self.base.retrieve.call_count == 1 + + def test_result_object_identity_unchanged_on_passthrough(self): + result = _make_result([self.entity], []) + base = _base_mock([result]) + tr = TemporalGraphRetriever(base) + returned = tr.retrieve("q") + assert returned[0] is result + + +# --------------------------------------------------------------------------- +# TemporalGraphRetriever — temporal filtering +# --------------------------------------------------------------------------- + +class TestTemporalGraphRetrieverFiltering: + + def setup_method(self): + self.base = _base_mock() + + def _set(self, entities, relationships): + self.base.retrieve.return_value = [_make_result(entities, relationships)] + + # Entity validity + + def test_entity_expired_before_at_time_excluded(self): + e_current = _make_entity("current", _utc(2020, 1, 1), _utc(2025, 1, 1)) + e_expired = _make_entity("expired", _utc(2018, 1, 1), _utc(2020, 6, 1)) + self._set([e_current, e_expired], []) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + ids = {e["id"] for e in results[0].related_entities} + assert "current" in ids + assert "expired" not in ids + + def test_entity_not_yet_started_at_at_time_excluded(self): + e_future = _make_entity("future", _utc(2025, 1, 1)) + e_current = _make_entity("current", _utc(2020, 1, 1)) + self._set([e_future, e_current], []) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + ids = {e["id"] for e in results[0].related_entities} + assert "current" in ids + assert "future" not in ids + + def test_entity_with_no_temporal_bounds_always_included(self): + e_timeless = _make_entity("timeless") # no valid_from / valid_until + self._set([e_timeless], []) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert len(results[0].related_entities) == 1 + + def test_entity_valid_on_boundary_date_included(self): + # valid_from == at_time — boundary should be inclusive + at = _utc(2022, 6, 1) + e = _make_entity("boundary", valid_from=at) + self._set([e], []) + + results = TemporalGraphRetriever(self.base, at_time=at).retrieve("q") + assert len(results[0].related_entities) == 1 + + def test_all_entities_expired_leaves_empty_list(self): + entities = [ + _make_entity("a", _utc(2010, 1, 1), _utc(2015, 1, 1)), + _make_entity("b", _utc(2012, 1, 1), _utc(2014, 1, 1)), + ] + self._set(entities, []) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert results[0].related_entities == [] + + # Relationship filtering + + def test_dangling_relationship_removed_when_target_expired(self): + e_a = _make_entity("A", _utc(2020, 1, 1)) + e_b = _make_entity("B", _utc(2022, 1, 1), _utc(2022, 6, 1)) + rel = _make_rel("A", "B") + self._set([e_a, e_b], [rel]) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert results[0].related_relationships == [] + + def test_dangling_relationship_removed_when_source_expired(self): + e_a = _make_entity("A", _utc(2020, 1, 1), _utc(2021, 1, 1)) + e_b = _make_entity("B", _utc(2020, 1, 1)) + rel = _make_rel("A", "B") + self._set([e_a, e_b], [rel]) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert results[0].related_relationships == [] + + def test_valid_relationship_kept(self): + e_a = _make_entity("A", _utc(2020, 1, 1)) + e_b = _make_entity("B", _utc(2020, 1, 1)) + rel = _make_rel("A", "B", valid_from=_utc(2020, 1, 1)) + self._set([e_a, e_b], [rel]) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert len(results[0].related_relationships) == 1 + + def test_relationship_with_expired_valid_until_removed(self): + e_a = _make_entity("A", _utc(2020, 1, 1)) + e_b = _make_entity("B", _utc(2020, 1, 1)) + rel = _make_rel("A", "B", valid_from=_utc(2020, 1, 1), valid_until=_utc(2021, 1, 1)) + self._set([e_a, e_b], [rel]) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert results[0].related_relationships == [] + + def test_multiple_relationship_types_filtered_independently(self): + e_a = _make_entity("A", _utc(2020, 1, 1)) + e_b = _make_entity("B", _utc(2020, 1, 1)) + e_c = _make_entity("C", _utc(2020, 1, 1), _utc(2021, 1, 1)) # expires + rel_ab = _make_rel("A", "B", rel_type="USES") + rel_ac = _make_rel("A", "C", rel_type="OWNS") + self._set([e_a, e_b, e_c], [rel_ab, rel_ac]) + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + rels = results[0].related_relationships + types = {r["type"] for r in rels} + assert "USES" in types + assert "OWNS" not in types + + # at_time precedence + + def test_call_site_at_time_overrides_constructor(self): + e = _make_entity("e", _utc(2018, 1, 1), _utc(2021, 1, 1)) + self._set([e], []) + + tr = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)) + results = tr.retrieve("q", at_time=_utc(2020, 6, 1)) + assert len(results[0].related_entities) == 1 + + def test_string_at_time_parsed_correctly(self): + e_current = _make_entity("current", _utc(2020, 1, 1)) + e_expired = _make_entity("expired", _utc(2018, 1, 1), _utc(2020, 6, 1)) + self._set([e_current, e_expired], []) + + results = TemporalGraphRetriever(self.base, at_time="2023-01-01").retrieve("q") + ids = {e["id"] for e in results[0].related_entities} + assert "current" in ids + assert "expired" not in ids + + def test_datetime_at_time_accepted_directly(self): + e = _make_entity("e", _utc(2020, 1, 1)) + self._set([e], []) + + results = TemporalGraphRetriever( + self.base, at_time=_utc(2023, 1, 1) + ).retrieve("q") + assert len(results[0].related_entities) == 1 + + # Multiple results + + def test_all_results_filtered_independently(self): + e_valid = _make_entity("valid", _utc(2020, 1, 1)) + e_old = _make_entity("old", _utc(2010, 1, 1), _utc(2015, 1, 1)) + r1 = _make_result([e_valid], []) + r2 = _make_result([e_old], []) + self.base.retrieve.return_value = [r1, r2] + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert len(results[0].related_entities) == 1 + assert len(results[1].related_entities) == 0 + + def test_result_scores_preserved_after_filtering(self): + e = _make_entity("e", _utc(2020, 1, 1)) + r = _make_result([e], [], score=0.77) + self.base.retrieve.return_value = [r] + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert results[0].score == pytest.approx(0.77) + + def test_result_content_preserved_after_filtering(self): + e = _make_entity("e", _utc(2020, 1, 1)) + r = _make_result([e], [], content="important drug interaction fact") + self.base.retrieve.return_value = [r] + + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert results[0].content == "important drug interaction fact" + + def test_empty_entities_and_relationships_stays_empty(self): + self.base.retrieve.return_value = [_make_result([], [])] + results = TemporalGraphRetriever(self.base, at_time=_utc(2023, 1, 1)).retrieve("q") + assert results[0].related_entities == [] + assert results[0].related_relationships == [] + + +# --------------------------------------------------------------------------- +# Temporal context header — _generate_reasoned_response +# --------------------------------------------------------------------------- + +class TestTemporalContextHeader: + + def setup_method(self): + self.retriever = ContextRetriever() + self.mock_llm = MagicMock() + self.mock_llm.generate.return_value = "LLM answer" + + def _call(self, at_time=None, header_template=None, contexts=None): + if contexts is None: + contexts = [RetrievedContext(content="fact A", score=0.9, source="graph")] + kwargs = {} + if header_template is not None: + kwargs["header_template"] = header_template + return self.retriever._generate_reasoned_response( + "test query", contexts, [], self.mock_llm, at_time=at_time, **kwargs + ) + + def _last_prompt(self): + return self.mock_llm.generate.call_args[0][0] + + def test_no_header_without_at_time(self): + self._call() + assert "Graph context valid as of" not in self._last_prompt() + + def test_header_present_when_at_time_datetime(self): + self._call(at_time=_utc(2023, 6, 1)) + prompt = self._last_prompt() + assert "2023-06-01" in prompt + assert "Graph context valid as of" in prompt + + def test_header_present_when_at_time_string(self): + self._call(at_time="2022-03-15") + assert "2022-03-15" in self._last_prompt() + + def test_header_appears_before_retrieved_context(self): + self._call(at_time=_utc(2023, 6, 1)) + prompt = self._last_prompt() + assert prompt.find("Graph context valid as of") < prompt.find("Retrieved Context:") + + def test_header_contains_source_label(self): + self._call(at_time=_utc(2023, 6, 1)) + assert "KnowledgeGraph snapshot" in self._last_prompt() + + def test_header_template_configurable(self): + self._call( + at_time=_utc(2023, 6, 1), + header_template="[Snapshot: {at_time} from {source}]", + ) + prompt = self._last_prompt() + assert "[Snapshot:" in prompt + assert "2023-06-01" in prompt + + def test_custom_template_source_placeholder_filled(self): + self._call( + at_time=_utc(2023, 1, 1), + header_template="Source={source}", + ) + assert "Source=KnowledgeGraph snapshot" in self._last_prompt() + + def test_prompt_contains_user_question(self): + self.retriever._generate_reasoned_response( + "How many suppliers?", [], [], self.mock_llm + ) + assert "How many suppliers?" in self._last_prompt() + + def test_prompt_contains_retrieved_context_content(self): + ctx = RetrievedContext(content="DrugA interaction warning", score=0.9, source="graph") + self.retriever._generate_reasoned_response( + "q", [ctx], [], self.mock_llm + ) + assert "DrugA interaction warning" in self._last_prompt() + + def test_no_at_time_prompt_identical_to_baseline(self): + ctx = RetrievedContext(content="fact", score=0.8, source="graph") + self.retriever._generate_reasoned_response("q", [ctx], [], self.mock_llm) + prompt_without = self._last_prompt() + + self.retriever._generate_reasoned_response("q", [ctx], [], self.mock_llm, at_time=None) + prompt_with_none = self._last_prompt() + + assert prompt_without == prompt_with_none + + def test_llm_generate_called_once_per_call(self): + self._call(at_time=_utc(2023, 1, 1)) + assert self.mock_llm.generate.call_count == 1 + + def test_query_with_reasoning_threads_at_time(self): + retriever = ContextRetriever() + retriever.retrieve = MagicMock(return_value=[]) + + with patch.object( + retriever, "_generate_reasoned_response", wraps=retriever._generate_reasoned_response + ) as mock_gen: + mock_gen.return_value = "answer" + retriever.query_with_reasoning( + "test query", self.mock_llm, at_time=_utc(2023, 1, 1) + ) + assert mock_gen.call_args.kwargs.get("at_time") == _utc(2023, 1, 1) + + def test_query_with_reasoning_threads_header_template(self): + retriever = ContextRetriever() + retriever.retrieve = MagicMock(return_value=[]) + custom = "[{at_time}|{source}]" + + with patch.object( + retriever, "_generate_reasoned_response", wraps=retriever._generate_reasoned_response + ) as mock_gen: + mock_gen.return_value = "answer" + retriever.query_with_reasoning( + "test query", self.mock_llm, + at_time=_utc(2023, 1, 1), + header_template=custom, + ) + assert mock_gen.call_args.kwargs.get("header_template") == custom + + def test_query_with_reasoning_no_at_time_no_header(self): + retriever = ContextRetriever() + retriever.retrieve = MagicMock(return_value=[ + RetrievedContext(content="fact", score=0.9, source="graph") + ]) + retriever.query_with_reasoning("q", self.mock_llm) + prompt = self.mock_llm.generate.call_args[0][0] + assert "Graph context valid as of" not in prompt diff --git a/tests/kg/test_temporal_query_rewriter.py b/tests/kg/test_temporal_query_rewriter.py new file mode 100644 index 00000000..eeac1265 --- /dev/null +++ b/tests/kg/test_temporal_query_rewriter.py @@ -0,0 +1,418 @@ +""" +Tests for TemporalQueryRewriter and TemporalQueryResult. +""" + +import json +from datetime import datetime, timezone +from unittest.mock import MagicMock, patch + +import pytest + +from semantica.kg import TemporalQueryRewriter, TemporalQueryResult + + +def _utc(year, month=1, day=1): + return datetime(year, month, day, tzinfo=timezone.utc) + + +# --------------------------------------------------------------------------- +# TemporalQueryResult — dataclass behaviour +# --------------------------------------------------------------------------- + +class TestTemporalQueryResult: + + def test_has_temporal_context_true_when_intent_set(self): + r = TemporalQueryResult( + rewritten_query="q", temporal_intent="before", + at_time=_utc(2021), confidence=0.9, + ) + assert r.has_temporal_context() is True + + def test_has_temporal_context_false_when_no_intent(self): + r = TemporalQueryResult(rewritten_query="q", confidence=1.0) + assert r.has_temporal_context() is False + + def test_default_fields_are_none(self): + r = TemporalQueryResult(rewritten_query="q", confidence=0.5) + assert r.at_time is None + assert r.start_time is None + assert r.end_time is None + assert r.temporal_intent is None + + def test_all_intent_values_accepted(self): + for intent in ("before", "after", "at", "during", "between"): + r = TemporalQueryResult(rewritten_query="q", temporal_intent=intent, confidence=0.8) + assert r.has_temporal_context() is True + + def test_between_populates_start_and_end(self): + r = TemporalQueryResult( + rewritten_query="q", + temporal_intent="between", + start_time=_utc(2019), + end_time=_utc(2022), + confidence=0.85, + ) + assert r.start_time < r.end_time + assert r.at_time is None + + +# --------------------------------------------------------------------------- +# No temporal phrase — passthrough +# --------------------------------------------------------------------------- + +class TestNoTemporalPhrase: + + def setup_method(self): + self.rw = TemporalQueryRewriter() + + def test_plain_query_rewritten_query_unchanged(self): + q = "what are the top suppliers?" + assert self.rw.rewrite(q).rewritten_query == q + + def test_plain_query_all_fields_none(self): + r = self.rw.rewrite("list certified partners") + assert r.temporal_intent is None + assert r.at_time is None + assert r.start_time is None + assert r.end_time is None + + def test_empty_string_returns_empty_rewritten(self): + r = self.rw.rewrite("") + assert r.rewritten_query == "" + assert r.temporal_intent is None + + def test_entity_name_with_year_not_matched_as_temporal(self): + # "ISO9001" or "GDPR2018" should not be treated as temporal + r = self.rw.rewrite("list all ISO9001 certified suppliers") + # May or may not extract — just must not crash + assert r.rewritten_query is not None + + def test_confidence_is_1_when_no_phrase_found(self): + r = self.rw.rewrite("active suppliers in good standing") + # No unambiguous temporal phrase — confidence should be high + assert r.confidence > 0.0 + + def test_no_temporal_phrase_never_sets_at_time(self): + for q in [ + "show all nodes", + "what is the status of project Alpha?", + "find entities of type Drug", + ]: + assert self.rw.rewrite(q).at_time is None + + +# --------------------------------------------------------------------------- +# "before" intent +# --------------------------------------------------------------------------- + +class TestBeforeIntent: + + def setup_method(self): + self.rw = TemporalQueryRewriter() + + def test_before_year_sets_intent(self): + r = self.rw.rewrite("which suppliers were certified before 2021?") + assert r.temporal_intent == "before" + + def test_before_year_sets_at_time(self): + r = self.rw.rewrite("which suppliers were certified before 2021?") + assert r.at_time is not None + assert r.at_time.year == 2021 + + def test_before_year_start_and_end_are_none(self): + r = self.rw.rewrite("approved before 2021") + assert r.start_time is None + assert r.end_time is None + + def test_prior_to_maps_to_before(self): + r = self.rw.rewrite("rules that applied prior to 2020") + assert r.temporal_intent == "before" + assert r.at_time.year == 2020 + + def test_until_maps_to_before(self): + r = self.rw.rewrite("approvals valid until 2022") + assert r.temporal_intent == "before" + + def test_before_strips_phrase_from_query(self): + r = self.rw.rewrite("which drugs were approved before 2019?") + assert "before" not in r.rewritten_query + assert "2019" not in r.rewritten_query + + def test_rewritten_query_not_empty_after_strip(self): + r = self.rw.rewrite("drugs approved before 2019") + assert r.rewritten_query.strip() != "" + + def test_before_different_years(self): + for year in (2010, 2015, 2020, 2023): + r = self.rw.rewrite(f"facts recorded before {year}") + assert r.at_time.year == year + + +# --------------------------------------------------------------------------- +# "after" intent +# --------------------------------------------------------------------------- + +class TestAfterIntent: + + def setup_method(self): + self.rw = TemporalQueryRewriter() + + def test_after_year_sets_intent(self): + r = self.rw.rewrite("which threat actors were active after 2018?") + assert r.temporal_intent == "after" + + def test_after_year_sets_at_time(self): + r = self.rw.rewrite("which threat actors were active after 2018?") + assert r.at_time is not None + assert r.at_time.year == 2018 + + def test_since_maps_to_after(self): + r = self.rw.rewrite("regulations introduced since 2015") + assert r.temporal_intent == "after" + + def test_following_maps_to_after(self): + r = self.rw.rewrite("changes following 2020") + assert r.temporal_intent == "after" + + def test_after_start_and_end_are_none(self): + r = self.rw.rewrite("events after 2018") + assert r.start_time is None + assert r.end_time is None + + def test_after_strips_phrase_from_query(self): + r = self.rw.rewrite("entities added after 2020") + assert "after" not in r.rewritten_query + + +# --------------------------------------------------------------------------- +# "during" / "in" intent +# --------------------------------------------------------------------------- + +class TestDuringIntent: + + def setup_method(self): + self.rw = TemporalQueryRewriter() + + def test_in_year_maps_to_during(self): + r = self.rw.rewrite("what interactions were known in 2021?") + assert r.temporal_intent == "during" + + def test_in_year_sets_at_time(self): + r = self.rw.rewrite("what interactions were known in 2021?") + assert r.at_time is not None + assert r.at_time.year == 2021 + + def test_during_keyword_sets_intent(self): + r = self.rw.rewrite("compliance status during Q2 2022") + assert r.temporal_intent == "during" + assert r.at_time is not None + + def test_within_maps_to_during(self): + r = self.rw.rewrite("certifications renewed within 2023") + assert r.temporal_intent == "during" + + def test_during_start_and_end_are_none(self): + r = self.rw.rewrite("events in 2022") + assert r.start_time is None + assert r.end_time is None + + +# --------------------------------------------------------------------------- +# "between" intent +# --------------------------------------------------------------------------- + +class TestBetweenIntent: + + def setup_method(self): + self.rw = TemporalQueryRewriter() + + def test_between_years_sets_intent(self): + r = self.rw.rewrite("contracts active between 2019 and 2022?") + assert r.temporal_intent == "between" + + def test_between_populates_start_and_end(self): + r = self.rw.rewrite("contracts active between 2019 and 2022?") + assert r.start_time is not None + assert r.end_time is not None + + def test_between_at_time_is_none(self): + r = self.rw.rewrite("contracts between 2019 and 2022") + assert r.at_time is None + + def test_between_start_before_end(self): + r = self.rw.rewrite("interactions between 2019 and 2022") + assert r.start_time < r.end_time + + def test_between_quarters(self): + r = self.rw.rewrite("revenue between Q1 2022 and Q3 2022") + assert r.temporal_intent == "between" + assert r.start_time is not None + assert r.end_time is not None + + def test_between_quarter_start_before_end(self): + r = self.rw.rewrite("revenue between Q1 2022 and Q3 2022") + assert r.start_time < r.end_time + + def test_between_strips_phrase(self): + r = self.rw.rewrite("contracts active between 2019 and 2022") + assert "between" not in r.rewritten_query + + def test_between_start_year_correct(self): + r = self.rw.rewrite("events between 2018 and 2021") + assert r.start_time.year == 2018 + + def test_between_end_year_correct(self): + r = self.rw.rewrite("events between 2018 and 2021") + assert r.end_time.year == 2021 + + +# --------------------------------------------------------------------------- +# "at" / "as of" intent +# --------------------------------------------------------------------------- + +class TestAtIntent: + + def setup_method(self): + self.rw = TemporalQueryRewriter() + + def test_as_of_sets_at_intent(self): + r = self.rw.rewrite("graph state as of 2023") + assert r.temporal_intent == "at" + + def test_as_of_sets_at_time(self): + r = self.rw.rewrite("graph state as of 2023") + assert r.at_time is not None + assert r.at_time.year == 2023 + + def test_as_of_start_and_end_none(self): + r = self.rw.rewrite("state as of 2022") + assert r.start_time is None + assert r.end_time is None + + +# --------------------------------------------------------------------------- +# Rewritten query quality +# --------------------------------------------------------------------------- + +class TestRewrittenQueryQuality: + + def setup_method(self): + self.rw = TemporalQueryRewriter() + + def test_no_double_spaces_after_strip(self): + r = self.rw.rewrite("suppliers certified before 2021 are valid") + assert " " not in r.rewritten_query + + def test_no_leading_trailing_whitespace(self): + r = self.rw.rewrite("before 2021 show suppliers") + assert r.rewritten_query == r.rewritten_query.strip() + + def test_question_mark_preserved(self): + r = self.rw.rewrite("which drugs were approved before 2019?") + assert r.rewritten_query.endswith("?") + + def test_non_temporal_words_preserved(self): + r = self.rw.rewrite("certified suppliers before 2021") + assert "certified" in r.rewritten_query + assert "suppliers" in r.rewritten_query + + +# --------------------------------------------------------------------------- +# Rewriter never calls reconstruct_at_time +# --------------------------------------------------------------------------- + +class TestNoReconstructAtTime: + + def test_reconstruct_at_time_never_called_on_regex_path(self): + rw = TemporalQueryRewriter() + with patch("semantica.kg.temporal_query.TemporalGraphQuery") as mock_tgq: + rw.rewrite("threat actors active before 2021") + mock_tgq.assert_not_called() + + def test_reconstruct_at_time_never_called_on_no_phrase(self): + rw = TemporalQueryRewriter() + with patch("semantica.kg.temporal_query.TemporalGraphQuery") as mock_tgq: + rw.rewrite("list all suppliers") + mock_tgq.assert_not_called() + + +# --------------------------------------------------------------------------- +# LLM-assisted path +# --------------------------------------------------------------------------- + +class TestLLMAssistedRewrite: + + def setup_method(self): + self.mock_llm = MagicMock() + + def _set_llm(self, phrase, intent, confidence=0.9): + self.mock_llm.generate.return_value = json.dumps({ + "temporal_phrase": phrase, + "temporal_intent": intent, + "confidence": confidence, + }) + + def test_llm_called_before_regex(self): + self._set_llm("2021", "before") + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + rw.rewrite("before 2021") + assert self.mock_llm.generate.call_count == 1 + + def test_llm_before_intent_sets_at_time(self): + self._set_llm("the 2021 merger", "before") + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + r = rw.rewrite("suppliers certified before the 2021 merger?") + assert r.temporal_intent == "before" + assert r.at_time is not None + + def test_llm_confidence_propagated(self): + self._set_llm("2022", "after", confidence=0.95) + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + r = rw.rewrite("events after 2022") + assert r.confidence == pytest.approx(0.95) + + def test_llm_between_intent_populates_range(self): + self._set_llm("Q1 and Q3 2022", "between") + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + r = rw.rewrite("revenue between Q1 and Q3 2022") + assert r.temporal_intent == "between" + + def test_llm_null_phrase_returns_no_temporal_context(self): + self.mock_llm.generate.return_value = json.dumps({ + "temporal_phrase": None, "temporal_intent": None, "confidence": 0.99, + }) + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + r = rw.rewrite("what are the top suppliers?") + assert r.temporal_intent is None + + def test_llm_failure_falls_back_to_regex(self): + self.mock_llm.generate.side_effect = RuntimeError("LLM offline") + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + r = rw.rewrite("contracts before 2020") + assert r.temporal_intent == "before" + assert r.at_time is not None + + def test_llm_invalid_json_falls_back_to_regex(self): + self.mock_llm.generate.return_value = "not json at all" + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + r = rw.rewrite("entities after 2019") + assert r.temporal_intent == "after" + + def test_llm_prompt_contains_query(self): + self._set_llm("2021", "before") + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + rw.rewrite("suppliers before 2021") + prompt = self.mock_llm.generate.call_args[0][0] + assert "suppliers before 2021" in prompt + + def test_llm_not_called_when_no_provider(self): + rw = TemporalQueryRewriter() # no llm_provider + rw.rewrite("before 2021") + self.mock_llm.generate.assert_not_called() + + def test_llm_rewrite_never_calls_reconstruct_at_time(self): + self._set_llm("2021", "before") + rw = TemporalQueryRewriter(llm_provider=self.mock_llm) + with patch("semantica.kg.temporal_query.TemporalGraphQuery") as mock_tgq: + rw.rewrite("suppliers before the 2021 audit") + mock_tgq.assert_not_called()