mirror of
https://github.com/semantica-agi/semantica.git
synced 2026-08-29 04:26:20 +00:00
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:
co-authored by
Claude Sonnet 4.6
parent
9e68266563
commit
0f0800f109
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user