feat(#402): Temporal GraphRAG Integration — TemporalGraphRetriever & TemporalQueryRewriter

- 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 <noreply@anthropic.com>
This commit is contained in:
KaifAhmad1
2026-03-26 14:20:38 +05:30
co-authored by Claude Sonnet 4.6
parent 9e68266563
commit 0f0800f109
6 changed files with 1517 additions and 7 deletions
+2 -1
View File
@@ -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",
+209 -6
View File
@@ -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
+3
View File
@@ -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
+457
View File
@@ -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<intent_kw>
between | prior\s+to | before | until | up\s+to |
after | since | following |
during | in | within |
as\s+of | at | on
)
\b
\s+
(?P<phrase>
# "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<start>[\w][\w\s\-,/]*?)\s+and\s+(?P<end>[\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
+428
View File
@@ -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
+418
View File
@@ -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()