876a2c6916
CI / benchmark-publish (push) Has started running
CI / lint (push) Successful in 55s
CI / quality (push) Successful in 1m6s
CI / typecheck (push) Successful in 1m27s
CI / helm (push) Successful in 31s
CI / push-validation (push) Successful in 32s
CI / security (push) Successful in 1m55s
CI / build (push) Successful in 49s
CI / benchmark-regression (push) Has been skipped
CI / integration_tests (push) Successful in 3m35s
CI / e2e_tests (push) Successful in 3m43s
CI / unit_tests (push) Successful in 4m34s
CI / docker (push) Successful in 1m28s
CI / coverage (push) Successful in 10m37s
CI / status-check (push) Successful in 3s
CI / benchmark-publish (pull_request) Has been skipped
CI / benchmark-regression (pull_request) Failing after 1m36s
CI / coverage (pull_request) Successful in 11m39s
CI / helm (pull_request) Successful in 45s
CI / lint (pull_request) Successful in 1m59s
CI / quality (pull_request) Successful in 2m10s
CI / typecheck (pull_request) Successful in 2m19s
CI / security (pull_request) Successful in 2m24s
CI / e2e_tests (pull_request) Successful in 4m56s
CI / integration_tests (pull_request) Successful in 5m15s
CI / unit_tests (pull_request) Successful in 6m56s
CI / docker (pull_request) Successful in 1m35s
CI / build (pull_request) Successful in 1m18s
CI / push-validation (pull_request) Successful in 42s
CI / status-check (pull_request) Successful in 3s
Implement dual-path session management in LLMTraceRepository.save(): - UoW mode (explicit session provided): flush only, caller controls commit - Standalone mode (no session): flush + commit + close for durable persistence This resolves three data-integrity violations: 1. Premature commit of outer UoW transactions 2. Loss of rollback capability for subsequent failures 3. Mismatch between class docstring and implementation Also adds: - Input validation: trace must not be None - Updated BDD step definitions to pass session explicitly in UoW scenarios - close() method to _BrokenSession mock for proper cleanup path coverage - CHANGELOG.md entry for issue #7505 - CONTRIBUTORS.md credit for HAL 9000 ISSUES CLOSED: #7505
1064 lines
36 KiB
Python
1064 lines
36 KiB
Python
"""Step definitions for the LLM trace observability feature."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import builtins
|
|
import os
|
|
from datetime import datetime
|
|
from typing import Any
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
from behave import given, then, when
|
|
from behave.runner import Context
|
|
from pydantic import ValidationError
|
|
from sqlalchemy import create_engine
|
|
from sqlalchemy.exc import DatabaseError as SQLAlchemyDatabaseError
|
|
from sqlalchemy.orm import sessionmaker
|
|
|
|
from cleveragents.application.services.trace_service import (
|
|
TraceService,
|
|
_forward_trace_to_langsmith,
|
|
)
|
|
from cleveragents.core.exceptions import DatabaseError as AppDatabaseError
|
|
from cleveragents.domain.models.observability.llm_trace import LLMTrace, LLMTraceQuery
|
|
from cleveragents.domain.models.observability.metrics import (
|
|
MetricCollector,
|
|
MetricEntry,
|
|
OperationalMetricKey,
|
|
)
|
|
from cleveragents.infrastructure.database.llm_trace_repository import LLMTraceRepository
|
|
from cleveragents.infrastructure.database.models import Base, LLMTraceModel
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Test constants
|
|
# ---------------------------------------------------------------------------
|
|
|
|
VALID_ULID_1 = "01HXAAAAAAAAAAAAAAAAAAAAAA"
|
|
VALID_ULID_2 = "01HXBBBBBBBBBBBBBBBBBBBBBB"
|
|
VALID_ULID_3 = "01HXCCCCCCCCCCCCCCCCCCCCCC"
|
|
VALID_PLAN_ID = "01HX0000000000PPPPPPPPPPPP"
|
|
VALID_DECISION_ID = "01HX0000000000DDDDDDDDDDDD"
|
|
|
|
|
|
def _make_trace(
|
|
trace_id: str = VALID_ULID_1,
|
|
plan_id: str = VALID_PLAN_ID,
|
|
decision_id: str | None = None,
|
|
actor: str = "planner",
|
|
provider: str = "openai",
|
|
model: str = "gpt-4o",
|
|
prompt_tokens: int = 500,
|
|
completion_tokens: int = 200,
|
|
cost_usd: float = 0.01,
|
|
latency_ms: float = 350.0,
|
|
**kwargs: Any,
|
|
) -> LLMTrace:
|
|
"""Build a test LLMTrace with sensible defaults."""
|
|
return LLMTrace(
|
|
trace_id=trace_id,
|
|
plan_id=plan_id,
|
|
decision_id=decision_id,
|
|
actor=actor,
|
|
provider=provider,
|
|
model=model,
|
|
prompt_tokens=prompt_tokens,
|
|
completion_tokens=completion_tokens,
|
|
cost_usd=cost_usd,
|
|
latency_ms=latency_ms,
|
|
**kwargs,
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# In-memory repository stub
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class InMemoryLLMTraceRepository:
|
|
"""Minimal in-memory repository for test isolation."""
|
|
|
|
def __init__(self) -> None:
|
|
self._store: dict[str, LLMTrace] = {}
|
|
|
|
def save(self, trace: LLMTrace) -> None:
|
|
self._store[trace.trace_id] = trace
|
|
|
|
def get(self, trace_id: str) -> LLMTrace | None:
|
|
return self._store.get(trace_id)
|
|
|
|
def list_by_plan(self, plan_id: str) -> list[LLMTrace]:
|
|
return sorted(
|
|
(t for t in self._store.values() if t.plan_id == plan_id),
|
|
key=lambda t: t.timestamp.isoformat(),
|
|
)
|
|
|
|
def list_by_decision(self, decision_id: str) -> list[LLMTrace]:
|
|
return sorted(
|
|
(t for t in self._store.values() if t.decision_id == decision_id),
|
|
key=lambda t: t.timestamp.isoformat(),
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Background
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("a trace service with an in-memory repository")
|
|
def step_create_service(context: Context) -> None:
|
|
context.repo = InMemoryLLMTraceRepository()
|
|
context.settings = MagicMock()
|
|
context.service = TraceService(
|
|
settings=context.settings,
|
|
repository=context.repo, # type: ignore[arg-type]
|
|
)
|
|
context.error = None
|
|
context.traces_list = [] # list[LLMTrace]
|
|
context.metrics_list = [] # list[MetricEntry]
|
|
context.langsmith_called = False
|
|
# Ensure LANGCHAIN_TRACING_V2 is unset by default
|
|
os.environ.pop("LANGCHAIN_TRACING_V2", None)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Model validation
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("a valid LLM trace payload")
|
|
def step_valid_payload(context: Context) -> None:
|
|
context.trace_kwargs = {
|
|
"trace_id": VALID_ULID_1,
|
|
"plan_id": VALID_PLAN_ID,
|
|
"actor": "planner",
|
|
"provider": "openai",
|
|
"model": "gpt-4o",
|
|
"prompt_tokens": 500,
|
|
"completion_tokens": 200,
|
|
"cost_usd": 0.01,
|
|
"latency_ms": 350.0,
|
|
}
|
|
|
|
|
|
@given("a minimal LLM trace payload")
|
|
def step_minimal_payload(context: Context) -> None:
|
|
context.trace_kwargs = {
|
|
"trace_id": VALID_ULID_1,
|
|
"plan_id": VALID_PLAN_ID,
|
|
"actor": "planner",
|
|
"provider": "openai",
|
|
"model": "gpt-4o",
|
|
"prompt_tokens": 0,
|
|
"completion_tokens": 0,
|
|
"cost_usd": 0.0,
|
|
"latency_ms": 0.0,
|
|
}
|
|
|
|
|
|
@when("I create the LLM trace model")
|
|
def step_create_model(context: Context) -> None:
|
|
try:
|
|
context.trace = LLMTrace(**context.trace_kwargs)
|
|
except ValidationError as exc:
|
|
context.error = exc
|
|
|
|
|
|
@when("I create a trace with an invalid trace_id")
|
|
def step_invalid_trace_id(context: Context) -> None:
|
|
try:
|
|
_make_trace(trace_id="INVALID")
|
|
except ValidationError as exc:
|
|
context.error = exc
|
|
|
|
|
|
@when("I create a trace with negative prompt_tokens")
|
|
def step_negative_prompt(context: Context) -> None:
|
|
try:
|
|
_make_trace(prompt_tokens=-1)
|
|
except ValidationError as exc:
|
|
context.error = exc
|
|
|
|
|
|
@when("I create a trace with negative cost_usd")
|
|
def step_negative_cost(context: Context) -> None:
|
|
try:
|
|
_make_trace(cost_usd=-0.5)
|
|
except ValidationError as exc:
|
|
context.error = exc
|
|
|
|
|
|
@then("the trace should have all required fields populated")
|
|
def step_all_fields(context: Context) -> None:
|
|
t = context.trace
|
|
assert t.trace_id == VALID_ULID_1
|
|
assert t.plan_id == VALID_PLAN_ID
|
|
assert t.actor == "planner"
|
|
assert t.provider == "openai"
|
|
assert t.model == "gpt-4o"
|
|
assert t.prompt_tokens == 500
|
|
assert t.completion_tokens == 200
|
|
|
|
|
|
@then("the trace should be frozen")
|
|
def step_frozen(context: Context) -> None:
|
|
try:
|
|
context.trace.actor = "modified" # type: ignore[misc]
|
|
raise AssertionError("Expected frozen model to reject assignment")
|
|
except ValidationError:
|
|
pass
|
|
|
|
|
|
@then("the trace decision_id should be None")
|
|
def step_decision_none(context: Context) -> None:
|
|
assert context.trace.decision_id is None
|
|
|
|
|
|
@then("the trace tool_calls should be empty")
|
|
def step_tool_calls_empty(context: Context) -> None:
|
|
assert context.trace.tool_calls == []
|
|
|
|
|
|
@then("the trace context_hash should be None")
|
|
def step_context_hash_none(context: Context) -> None:
|
|
assert context.trace.context_hash is None
|
|
|
|
|
|
@then("the trace streaming should be False")
|
|
def step_streaming_false(context: Context) -> None:
|
|
assert context.trace.streaming is False
|
|
|
|
|
|
@then("the trace retry_count should be 0")
|
|
def step_retry_zero(context: Context) -> None:
|
|
assert context.trace.retry_count == 0
|
|
|
|
|
|
@then("the trace error should be None")
|
|
def step_error_none(context: Context) -> None:
|
|
assert context.trace.error is None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Query model
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I create a trace query with plan_id filter")
|
|
def step_create_query(context: Context) -> None:
|
|
context.query = LLMTraceQuery(plan_id=VALID_PLAN_ID)
|
|
|
|
|
|
@when("I create a trace query with limit {limit:d}")
|
|
def step_query_bad_limit(context: Context, limit: int) -> None:
|
|
try:
|
|
LLMTraceQuery(limit=limit)
|
|
except ValidationError as exc:
|
|
context.error = exc
|
|
|
|
|
|
@then("the query should have default limit {limit:d}")
|
|
def step_query_limit(context: Context, limit: int) -> None:
|
|
assert context.query.limit == limit
|
|
|
|
|
|
@then("the query should have default offset {offset:d}")
|
|
def step_query_offset(context: Context, offset: int) -> None:
|
|
assert context.query.offset == offset
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Metric keys
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@then("the OperationalMetricKey enum should have exactly {count:d} members")
|
|
def step_metric_key_count(context: Context, count: int) -> None:
|
|
assert len(OperationalMetricKey) == count
|
|
|
|
|
|
@then('PLAN_DURATION_MS should equal "{value}"')
|
|
def step_plan_duration_value(context: Context, value: str) -> None:
|
|
assert OperationalMetricKey.PLAN_DURATION_MS.value == value
|
|
|
|
|
|
@then('LLM_AVG_LATENCY_MS should equal "{value}"')
|
|
def step_llm_avg_latency_value(context: Context, value: str) -> None:
|
|
assert OperationalMetricKey.LLM_AVG_LATENCY_MS.value == value
|
|
|
|
|
|
@then('SUBPLAN_COUNT should equal "{value}"')
|
|
def step_subplan_count_value(context: Context, value: str) -> None:
|
|
assert OperationalMetricKey.SUBPLAN_COUNT.value == value
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Metric entry
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("a plan_id for metrics")
|
|
def step_plan_id_for_metrics(context: Context) -> None:
|
|
context.metric_plan_id = VALID_PLAN_ID
|
|
|
|
|
|
@when("I record a PLAN_DURATION_MS metric with value {value:g}")
|
|
def step_record_metric(context: Context, value: float) -> None:
|
|
context.metric_entry = MetricCollector.record(
|
|
OperationalMetricKey.PLAN_DURATION_MS,
|
|
value,
|
|
context.metric_plan_id,
|
|
)
|
|
|
|
|
|
@when("I record an ACTOR_LATENCY_MS metric with actor label")
|
|
def step_record_metric_with_label(context: Context) -> None:
|
|
context.metric_entry = MetricCollector.record(
|
|
OperationalMetricKey.ACTOR_LATENCY_MS,
|
|
150.0,
|
|
context.metric_plan_id,
|
|
labels={"actor": "planner"},
|
|
)
|
|
|
|
|
|
@then("the metric entry key should be PLAN_DURATION_MS")
|
|
def step_metric_key_plan_duration(context: Context) -> None:
|
|
assert context.metric_entry.key == OperationalMetricKey.PLAN_DURATION_MS
|
|
|
|
|
|
@then("the metric entry value should be {value:g}")
|
|
def step_metric_value(context: Context, value: float) -> None:
|
|
assert context.metric_entry.value == value
|
|
|
|
|
|
@then("the metric entry should have a timestamp")
|
|
def step_metric_timestamp(context: Context) -> None:
|
|
assert isinstance(context.metric_entry.timestamp, datetime)
|
|
|
|
|
|
@then("the metric entry labels should contain actor")
|
|
def step_metric_labels(context: Context) -> None:
|
|
assert "actor" in context.metric_entry.labels
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Trace recording
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("a valid LLM trace")
|
|
def step_valid_trace(context: Context) -> None:
|
|
context.trace = _make_trace()
|
|
|
|
|
|
@given("three valid LLM traces for the same plan")
|
|
def step_three_traces(context: Context) -> None:
|
|
context.traces_list = [
|
|
_make_trace(
|
|
trace_id=VALID_ULID_1, prompt_tokens=500, cost_usd=0.01, latency_ms=300.0
|
|
),
|
|
_make_trace(
|
|
trace_id=VALID_ULID_2, prompt_tokens=300, cost_usd=0.005, latency_ms=200.0
|
|
),
|
|
_make_trace(
|
|
trace_id=VALID_ULID_3, prompt_tokens=700, cost_usd=0.02, latency_ms=500.0
|
|
),
|
|
]
|
|
|
|
|
|
@given("two traces with the same decision_id")
|
|
def step_two_traces_same_decision(context: Context) -> None:
|
|
context.traces_list = [
|
|
_make_trace(trace_id=VALID_ULID_1, decision_id=VALID_DECISION_ID),
|
|
_make_trace(trace_id=VALID_ULID_2, decision_id=VALID_DECISION_ID),
|
|
]
|
|
|
|
|
|
@when("I record the trace via the service")
|
|
def step_record_trace(context: Context) -> None:
|
|
def mock_internal_forward(trace: LLMTrace) -> None:
|
|
context.langsmith_called = True
|
|
if getattr(context, "langsmith_will_fail", False):
|
|
raise RuntimeError("LangSmith connection failed")
|
|
|
|
try:
|
|
with patch(
|
|
"cleveragents.application.services.trace_service._forward_trace_to_langsmith",
|
|
side_effect=mock_internal_forward,
|
|
):
|
|
context.service.record_trace(context.trace)
|
|
except Exception as exc:
|
|
context.error = exc
|
|
|
|
|
|
@when("I record all traces via the service")
|
|
def step_record_all_traces(context: Context) -> None:
|
|
for t in context.traces_list:
|
|
context.service.record_trace(t)
|
|
|
|
|
|
@then("I should be able to retrieve it by trace_id")
|
|
def step_retrieve_by_id(context: Context) -> None:
|
|
result = context.service.get_trace(context.trace.trace_id)
|
|
assert result is not None
|
|
assert result.trace_id == context.trace.trace_id
|
|
|
|
|
|
@then("I should be able to list it by plan_id")
|
|
def step_list_by_plan(context: Context) -> None:
|
|
results = context.service.get_traces(context.trace.plan_id)
|
|
assert len(results) >= 1
|
|
assert any(t.trace_id == context.trace.trace_id for t in results)
|
|
|
|
|
|
@then("listing by plan_id should return {count:d} traces")
|
|
def step_list_count(context: Context, count: int) -> None:
|
|
plan_id = context.traces_list[0].plan_id
|
|
results = context.service.get_traces(plan_id)
|
|
assert len(results) == count
|
|
|
|
|
|
@then("listing by decision_id should return {count:d} traces")
|
|
def step_list_by_decision_count(context: Context, count: int) -> None:
|
|
decision_id = context.traces_list[0].decision_id
|
|
assert decision_id is not None
|
|
results = context.service.get_traces_by_decision(decision_id)
|
|
assert len(results) == count
|
|
|
|
|
|
@when("I query a non-existent trace_id")
|
|
def step_query_missing(context: Context) -> None:
|
|
context.result = context.service.get_trace("01HXZZZZZZZZZZZZZZZZZZZZZZ")
|
|
|
|
|
|
@then("the trace query result should be None")
|
|
def step_result_none(context: Context) -> None:
|
|
assert context.result is None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Metric computation
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I compute metrics for the plan")
|
|
def step_compute_metrics(context: Context) -> None:
|
|
plan_id = context.traces_list[0].plan_id
|
|
context.metrics_list = context.service.compute_metrics(plan_id)
|
|
|
|
|
|
@when("I compute metrics for a plan with no traces")
|
|
def step_compute_empty(context: Context) -> None:
|
|
context.metrics_list = context.service.compute_metrics("01HXEEEEEEEEEEEEEEEEEEEEEE")
|
|
|
|
|
|
@then("I should get LLM_CALL_COUNT equal to {count:d}")
|
|
def step_llm_call_count(context: Context, count: int) -> None:
|
|
entry = _find_metric(context.metrics_list, OperationalMetricKey.LLM_CALL_COUNT)
|
|
assert entry is not None
|
|
assert entry.value == float(count)
|
|
|
|
|
|
@then("I should get LLM_TOTAL_TOKENS greater than {value:d}")
|
|
def step_llm_total_tokens(context: Context, value: int) -> None:
|
|
entry = _find_metric(context.metrics_list, OperationalMetricKey.LLM_TOTAL_TOKENS)
|
|
assert entry is not None
|
|
assert entry.value > value
|
|
|
|
|
|
@then("I should get LLM_TOTAL_COST_USD greater than {value:d}")
|
|
def step_llm_total_cost(context: Context, value: int) -> None:
|
|
entry = _find_metric(context.metrics_list, OperationalMetricKey.LLM_TOTAL_COST_USD)
|
|
assert entry is not None
|
|
assert entry.value > value
|
|
|
|
|
|
@then("I should get LLM_AVG_LATENCY_MS greater than {value:d}")
|
|
def step_llm_avg_latency(context: Context, value: int) -> None:
|
|
entry = _find_metric(context.metrics_list, OperationalMetricKey.LLM_AVG_LATENCY_MS)
|
|
assert entry is not None
|
|
assert entry.value > value
|
|
|
|
|
|
@then("the metrics list should be empty")
|
|
def step_metrics_empty(context: Context) -> None:
|
|
assert len(context.metrics_list) == 0
|
|
|
|
|
|
def _find_metric(
|
|
metrics: list[MetricEntry], key: OperationalMetricKey
|
|
) -> MetricEntry | None:
|
|
for m in metrics:
|
|
if m.key == key:
|
|
return m
|
|
return None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Lifecycle hooks
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I call on_plan_start")
|
|
def step_on_plan_start(context: Context) -> None:
|
|
context.metric_entry = context.service.on_plan_start(context.metric_plan_id)
|
|
|
|
|
|
@then("I should get a PLAN_DECISION_COUNT metric with value {value:d}")
|
|
def step_plan_decision_count(context: Context, value: int) -> None:
|
|
assert context.metric_entry is not None
|
|
assert context.metric_entry.key == OperationalMetricKey.PLAN_DECISION_COUNT
|
|
assert context.metric_entry.value == float(value)
|
|
|
|
|
|
@when('I call on_actor_invocation with actor "{actor}" and latency {latency:g}')
|
|
def step_on_actor_invocation(context: Context, actor: str, latency: float) -> None:
|
|
context.hook_metrics = context.service.on_actor_invocation(
|
|
context.metric_plan_id, actor, latency
|
|
)
|
|
|
|
|
|
@then("I should get ACTOR_INVOCATION_COUNT and ACTOR_LATENCY_MS metrics")
|
|
def step_actor_metrics(context: Context) -> None:
|
|
keys = {m.key for m in context.hook_metrics}
|
|
assert OperationalMetricKey.ACTOR_INVOCATION_COUNT in keys
|
|
assert OperationalMetricKey.ACTOR_LATENCY_MS in keys
|
|
|
|
|
|
@when('I call on_tool_execution with tool "{tool}" and no error')
|
|
def step_on_tool_no_error(context: Context, tool: str) -> None:
|
|
context.hook_metrics = context.service.on_tool_execution(
|
|
context.metric_plan_id, tool, errored=False
|
|
)
|
|
|
|
|
|
@when('I call on_tool_execution with tool "{tool}" and error')
|
|
def step_on_tool_error(context: Context, tool: str) -> None:
|
|
context.hook_metrics = context.service.on_tool_execution(
|
|
context.metric_plan_id, tool, errored=True
|
|
)
|
|
|
|
|
|
@then("I should get a TOOL_INVOCATION_COUNT metric")
|
|
def step_tool_invocation(context: Context) -> None:
|
|
keys = {m.key for m in context.hook_metrics}
|
|
assert OperationalMetricKey.TOOL_INVOCATION_COUNT in keys
|
|
|
|
|
|
@then("I should not get a TOOL_ERROR_RATE metric")
|
|
def step_no_tool_error_rate(context: Context) -> None:
|
|
keys = {m.key for m in context.hook_metrics}
|
|
assert OperationalMetricKey.TOOL_ERROR_RATE not in keys
|
|
|
|
|
|
@then("I should get both TOOL_INVOCATION_COUNT and TOOL_ERROR_RATE metrics")
|
|
def step_tool_both_metrics(context: Context) -> None:
|
|
keys = {m.key for m in context.hook_metrics}
|
|
assert OperationalMetricKey.TOOL_INVOCATION_COUNT in keys
|
|
assert OperationalMetricKey.TOOL_ERROR_RATE in keys
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# LangSmith forwarding
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given('LANGCHAIN_TRACING_V2 is set to "{value}"')
|
|
def step_langsmith_enabled(context: Context, value: str) -> None:
|
|
os.environ["LANGCHAIN_TRACING_V2"] = value
|
|
|
|
|
|
@given("LANGCHAIN_TRACING_V2 is not set")
|
|
def step_langsmith_disabled(context: Context) -> None:
|
|
os.environ.pop("LANGCHAIN_TRACING_V2", None)
|
|
|
|
|
|
@given("the LangSmith forwarder will raise an exception")
|
|
def step_langsmith_error(context: Context) -> None:
|
|
context.langsmith_will_fail = True
|
|
|
|
|
|
@then("the LangSmith forwarder should have been called")
|
|
def step_langsmith_called(context: Context) -> None:
|
|
assert context.langsmith_called is True
|
|
|
|
|
|
@then("the LangSmith forwarder should not have been called")
|
|
def step_langsmith_not_called(context: Context) -> None:
|
|
assert context.langsmith_called is False
|
|
|
|
|
|
@then("no exception should propagate")
|
|
def step_no_exception(context: Context) -> None:
|
|
assert context.error is None
|
|
|
|
|
|
@then("the trace should still be persisted")
|
|
def step_trace_persisted(context: Context) -> None:
|
|
result = context.service.get_trace(context.trace.trace_id)
|
|
assert result is not None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Convenience collectors
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I call MetricCollector.plan_duration with {value:g}")
|
|
def step_collector_plan_duration(context: Context, value: float) -> None:
|
|
context.metric_entry = MetricCollector.plan_duration(context.metric_plan_id, value)
|
|
|
|
|
|
@when("I call MetricCollector.plan_cost with {value:g}")
|
|
def step_collector_plan_cost(context: Context, value: float) -> None:
|
|
context.metric_entry = MetricCollector.plan_cost(context.metric_plan_id, value)
|
|
|
|
|
|
@when("I call MetricCollector.llm_call_count with {value:d}")
|
|
def step_collector_llm_call_count(context: Context, value: int) -> None:
|
|
context.metric_entry = MetricCollector.llm_call_count(
|
|
context.metric_plan_id, float(value)
|
|
)
|
|
|
|
|
|
@then("the metric key should be PLAN_DURATION_MS")
|
|
def step_key_plan_duration(context: Context) -> None:
|
|
assert context.metric_entry.key == OperationalMetricKey.PLAN_DURATION_MS
|
|
|
|
|
|
@then("the metric key should be PLAN_TOTAL_COST_USD")
|
|
def step_key_plan_cost(context: Context) -> None:
|
|
assert context.metric_entry.key == OperationalMetricKey.PLAN_TOTAL_COST_USD
|
|
|
|
|
|
@then("the metric key should be LLM_CALL_COUNT")
|
|
def step_key_llm_call(context: Context) -> None:
|
|
assert context.metric_entry.key == OperationalMetricKey.LLM_CALL_COUNT
|
|
|
|
|
|
@then("the metric value should be {value:g}")
|
|
def step_metric_exact_value(context: Context, value: float) -> None:
|
|
assert context.metric_entry.value == value
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Database model
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@then('the LLMTraceModel should have tablename "{name}"')
|
|
def step_table_name(context: Context, name: str) -> None:
|
|
assert LLMTraceModel.__tablename__ == name
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Repository integration (SQLAlchemy in-memory)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _create_in_memory_engine_and_session():
|
|
"""Create an in-memory SQLite engine and session factory for testing."""
|
|
engine = create_engine("sqlite:///:memory:")
|
|
Base.metadata.create_all(engine)
|
|
factory = sessionmaker(bind=engine)
|
|
return engine, factory
|
|
|
|
|
|
@when("I create a repository with None session_factory")
|
|
def step_repo_none_factory(context: Context) -> None:
|
|
try:
|
|
LLMTraceRepository(session_factory=None) # type: ignore[arg-type]
|
|
except ValueError as exc:
|
|
context.error = exc
|
|
|
|
|
|
@then("a ValueError should be raised for the repository")
|
|
def step_repo_valueerror_raised(context: Context) -> None:
|
|
assert context.error is not None, "Expected a ValueError"
|
|
assert isinstance(context.error, ValueError)
|
|
|
|
|
|
@given("a SQLAlchemy in-memory repository")
|
|
def step_sqla_repo(context: Context) -> None:
|
|
engine, factory = _create_in_memory_engine_and_session()
|
|
context.sqla_engine = engine
|
|
context.sqla_factory = factory
|
|
# Use a single session for all repository operations so that
|
|
# flush() (without commit) keeps data visible across calls.
|
|
shared_session = factory()
|
|
context.sqla_shared_session = shared_session
|
|
context.sqla_repo = LLMTraceRepository(session_factory=lambda: shared_session)
|
|
|
|
|
|
@given("a valid LLM trace with tool calls")
|
|
def step_trace_with_tool_calls(context: Context) -> None:
|
|
context.trace = _make_trace(
|
|
tool_calls=[{"name": "search", "arguments": {"q": "test"}}],
|
|
)
|
|
|
|
|
|
@when("I save the trace via the repository")
|
|
def step_save_via_repo(context: Context) -> None:
|
|
context.sqla_repo.save(context.trace)
|
|
|
|
|
|
@then("I should be able to get it by trace_id from the repository")
|
|
def step_get_from_repo(context: Context) -> None:
|
|
result = context.sqla_repo.get(context.trace.trace_id)
|
|
assert result is not None
|
|
assert result.trace_id == context.trace.trace_id
|
|
assert result.plan_id == context.trace.plan_id
|
|
assert result.actor == context.trace.actor
|
|
assert result.provider == context.trace.provider
|
|
assert result.model == context.trace.model
|
|
assert result.prompt_tokens == context.trace.prompt_tokens
|
|
assert result.completion_tokens == context.trace.completion_tokens
|
|
context.repo_result = result
|
|
|
|
|
|
@then("the retrieved trace should have tool calls")
|
|
def step_retrieved_has_tool_calls(context: Context) -> None:
|
|
result = context.sqla_repo.get(context.trace.trace_id)
|
|
assert result is not None
|
|
assert len(result.tool_calls) > 0
|
|
assert result.tool_calls[0]["name"] == "search"
|
|
|
|
|
|
@when("I get a non-existent trace from the repository")
|
|
def step_get_nonexistent_from_repo(context: Context) -> None:
|
|
context.repo_result = context.sqla_repo.get("01HXZZZZZZZZZZZZZZZZZZZZZZ")
|
|
|
|
|
|
@then("the repository result should be None")
|
|
def step_repo_result_none(context: Context) -> None:
|
|
assert context.repo_result is None
|
|
|
|
|
|
@given("two traces for the same plan in the repository")
|
|
def step_two_traces_same_plan_repo(context: Context) -> None:
|
|
context.repo_traces = [
|
|
_make_trace(trace_id=VALID_ULID_1, plan_id=VALID_PLAN_ID),
|
|
_make_trace(trace_id=VALID_ULID_2, plan_id=VALID_PLAN_ID),
|
|
]
|
|
|
|
|
|
@when("I save both traces via the repository")
|
|
def step_save_both_repo(context: Context) -> None:
|
|
for t in context.repo_traces:
|
|
context.sqla_repo.save(t)
|
|
|
|
|
|
@then("listing by plan from the repository should return {count:d} traces")
|
|
def step_list_by_plan_repo(context: Context, count: int) -> None:
|
|
results = context.sqla_repo.list_by_plan(VALID_PLAN_ID)
|
|
assert len(results) == count
|
|
|
|
|
|
@when("I list traces for a non-existent plan from the repository")
|
|
def step_list_nonexistent_plan_repo(context: Context) -> None:
|
|
context.repo_list_result = context.sqla_repo.list_by_plan(
|
|
"01HXZZZZZZZZZZZZZZZZZZZZZZ"
|
|
)
|
|
|
|
|
|
@then("the repository list should be empty")
|
|
def step_repo_list_empty(context: Context) -> None:
|
|
assert len(context.repo_list_result) == 0
|
|
|
|
|
|
@given("two traces for the same decision in the repository")
|
|
def step_two_traces_same_decision_repo(context: Context) -> None:
|
|
context.repo_traces = [
|
|
_make_trace(
|
|
trace_id=VALID_ULID_1,
|
|
plan_id=VALID_PLAN_ID,
|
|
decision_id=VALID_DECISION_ID,
|
|
),
|
|
_make_trace(
|
|
trace_id=VALID_ULID_2,
|
|
plan_id=VALID_PLAN_ID,
|
|
decision_id=VALID_DECISION_ID,
|
|
),
|
|
]
|
|
|
|
|
|
@then("listing by decision from the repository should return {count:d} traces")
|
|
def step_list_by_decision_repo(context: Context, count: int) -> None:
|
|
results = context.sqla_repo.list_by_decision(VALID_DECISION_ID)
|
|
assert len(results) == count
|
|
|
|
|
|
@when("I list traces for a non-existent decision from the repository")
|
|
def step_list_nonexistent_decision_repo(context: Context) -> None:
|
|
context.repo_list_result = context.sqla_repo.list_by_decision(
|
|
"01HXZZZZZZZZZZZZZZZZZZZZZZ"
|
|
)
|
|
|
|
|
|
# --- Repository error paths -----------------------------------------------
|
|
|
|
|
|
class _BrokenSession:
|
|
"""A mock session that raises SQLAlchemyDatabaseError on all ops."""
|
|
|
|
def add(self, _obj: Any) -> None:
|
|
raise SQLAlchemyDatabaseError("mock", {}, Exception("broken"))
|
|
|
|
def commit(self) -> None:
|
|
raise SQLAlchemyDatabaseError("mock", {}, Exception("broken"))
|
|
|
|
def rollback(self) -> None:
|
|
pass
|
|
|
|
def close(self) -> None:
|
|
pass
|
|
|
|
def query(self, *_args: Any, **_kwargs: Any) -> Any:
|
|
raise SQLAlchemyDatabaseError("mock", {}, Exception("broken"))
|
|
|
|
|
|
@given("a SQLAlchemy repository with a broken session")
|
|
def step_broken_repo(context: Context) -> None:
|
|
context.broken_repo = LLMTraceRepository(
|
|
session_factory=lambda: _BrokenSession() # type: ignore[return-value]
|
|
)
|
|
|
|
|
|
@when("I save the trace via the broken repository")
|
|
def step_save_broken(context: Context) -> None:
|
|
try:
|
|
context.broken_repo.save(context.trace)
|
|
except AppDatabaseError as exc:
|
|
context.error = exc
|
|
except Exception as exc:
|
|
# tenacity may wrap in RetryError; unwrap to the root cause
|
|
cause = getattr(exc, "__cause__", None) or exc
|
|
context.error = cause
|
|
|
|
|
|
@when("I get a trace from the broken repository")
|
|
def step_get_broken(context: Context) -> None:
|
|
try:
|
|
context.broken_repo.get(VALID_ULID_1)
|
|
except AppDatabaseError as exc:
|
|
context.error = exc
|
|
except Exception as exc:
|
|
cause = getattr(exc, "__cause__", None) or exc
|
|
context.error = cause
|
|
|
|
|
|
@when("I list by plan from the broken repository")
|
|
def step_list_plan_broken(context: Context) -> None:
|
|
try:
|
|
context.broken_repo.list_by_plan(VALID_PLAN_ID)
|
|
except AppDatabaseError as exc:
|
|
context.error = exc
|
|
except Exception as exc:
|
|
cause = getattr(exc, "__cause__", None) or exc
|
|
context.error = cause
|
|
|
|
|
|
@when("I list by decision from the broken repository")
|
|
def step_list_decision_broken(context: Context) -> None:
|
|
try:
|
|
context.broken_repo.list_by_decision(VALID_DECISION_ID)
|
|
except AppDatabaseError as exc:
|
|
context.error = exc
|
|
except Exception as exc:
|
|
cause = getattr(exc, "__cause__", None) or exc
|
|
context.error = cause
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# LangSmith internal forwarder detail tests
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("the langsmith SDK is mocked as available")
|
|
def step_langsmith_sdk_available(context: Context) -> None:
|
|
context.mock_langsmith_client = MagicMock()
|
|
context.mock_langsmith_module = MagicMock()
|
|
context.mock_langsmith_module.Client.return_value = context.mock_langsmith_client
|
|
|
|
|
|
@given('a valid LLM trace with error "{error_msg}"')
|
|
def step_trace_with_error(context: Context, error_msg: str) -> None:
|
|
context.trace = _make_trace(error=error_msg)
|
|
|
|
|
|
@when("I call the internal LangSmith forwarder")
|
|
def step_call_internal_forwarder(context: Context) -> None:
|
|
if hasattr(context, "mock_langsmith_module"):
|
|
# Mock the langsmith import inside _forward_trace_to_langsmith
|
|
real_import = builtins.__import__
|
|
|
|
def mock_import(name: str, *args: Any, **kwargs: Any) -> Any:
|
|
if name == "langsmith":
|
|
return context.mock_langsmith_module
|
|
return real_import(name, *args, **kwargs)
|
|
|
|
with patch.object(builtins, "__import__", side_effect=mock_import):
|
|
_forward_trace_to_langsmith(context.trace)
|
|
elif getattr(context, "langsmith_not_importable", False):
|
|
# Block the langsmith import to simulate SDK not installed
|
|
real_import = builtins.__import__
|
|
|
|
def block_langsmith(name: str, *args: Any, **kwargs: Any) -> Any:
|
|
if name == "langsmith":
|
|
raise ImportError("No module named 'langsmith'")
|
|
return real_import(name, *args, **kwargs)
|
|
|
|
with patch.object(builtins, "__import__", side_effect=block_langsmith):
|
|
_forward_trace_to_langsmith(context.trace)
|
|
else:
|
|
_forward_trace_to_langsmith(context.trace)
|
|
|
|
|
|
@then("the langsmith Client create_run should have been called")
|
|
def step_langsmith_create_run_called(context: Context) -> None:
|
|
context.mock_langsmith_client.create_run.assert_called_once()
|
|
|
|
|
|
@then("the langsmith Client create_run should have been called with error")
|
|
def step_langsmith_create_run_with_error(context: Context) -> None:
|
|
context.mock_langsmith_client.create_run.assert_called_once()
|
|
call_kwargs = context.mock_langsmith_client.create_run.call_args
|
|
# The call uses **run_data, so check kwargs
|
|
all_args = call_kwargs.kwargs if call_kwargs.kwargs else {}
|
|
assert "error" in all_args, f"Expected 'error' in call kwargs: {all_args}"
|
|
|
|
|
|
@given("the langsmith SDK is not importable")
|
|
def step_langsmith_not_importable(context: Context) -> None:
|
|
context.langsmith_not_importable = True
|
|
context.langsmith_client_created = False
|
|
|
|
|
|
@then("no langsmith Client should have been created")
|
|
def step_no_langsmith_client(context: Context) -> None:
|
|
# Verified by the fact that no error was raised and the function
|
|
# returned cleanly when the import was blocked.
|
|
assert not getattr(context, "langsmith_client_created", False)
|
|
|
|
|
|
@then("the langsmith_enabled check should return True")
|
|
def step_langsmith_enabled_true(context: Context) -> None:
|
|
assert TraceService._langsmith_enabled() is True
|
|
|
|
|
|
@then("the raw tool_calls_json in the database should be null")
|
|
def step_raw_tool_calls_null(context: Context) -> None:
|
|
session = getattr(context, "sqla_shared_session", None) or context.sqla_factory()
|
|
row = (
|
|
session.query(LLMTraceModel).filter_by(trace_id=context.trace.trace_id).first()
|
|
)
|
|
assert row is not None
|
|
assert row.tool_calls_json is None
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Session contract: flush not commit
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@when("I save the trace via the repository and spy on the session")
|
|
def step_save_with_spy(context: Context) -> None:
|
|
engine, factory = _create_in_memory_engine_and_session()
|
|
context.sqla_engine = engine
|
|
context.sqla_factory = factory
|
|
|
|
# Create a real session but wrap it to spy on flush/commit calls
|
|
real_session = factory()
|
|
context.spy_session = real_session
|
|
context.flush_called = False
|
|
context.commit_called = False
|
|
|
|
original_flush = real_session.flush
|
|
original_commit = real_session.commit
|
|
|
|
def spy_flush(*args: Any, **kwargs: Any) -> None:
|
|
context.flush_called = True
|
|
return original_flush(*args, **kwargs)
|
|
|
|
def spy_commit(*args: Any, **kwargs: Any) -> None:
|
|
context.commit_called = True
|
|
return original_commit(*args, **kwargs)
|
|
|
|
# Assign spy functions to session methods
|
|
# Using object.__setattr__ to bypass type checking for method replacement
|
|
object.__setattr__(real_session, "flush", spy_flush)
|
|
object.__setattr__(real_session, "commit", spy_commit)
|
|
|
|
# Pass the session explicitly to test the UoW path: save() must flush
|
|
# but must NOT commit (the caller owns the transaction boundary).
|
|
repo = LLMTraceRepository(session_factory=lambda: real_session)
|
|
repo.save(context.trace, session=real_session)
|
|
# Commit so the data is visible for subsequent queries
|
|
object.__setattr__(real_session, "commit", original_commit)
|
|
real_session.commit()
|
|
|
|
|
|
@then("the session flush should have been called")
|
|
def step_flush_called(context: Context) -> None:
|
|
assert context.flush_called is True, "Expected session.flush() to have been called"
|
|
|
|
|
|
@then("the session commit should not have been called")
|
|
def step_commit_not_called(context: Context) -> None:
|
|
assert context.commit_called is False, (
|
|
"Expected session.commit() NOT to have been called by save(), but it was called"
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# UnitOfWork rollback propagation
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@given("a UnitOfWork with an in-memory database")
|
|
def step_uow_in_memory(context: Context) -> None:
|
|
engine = create_engine(
|
|
"sqlite:///:memory:",
|
|
connect_args={"check_same_thread": False},
|
|
)
|
|
Base.metadata.create_all(engine)
|
|
context.uow_engine = engine
|
|
context.uow_session_factory = sessionmaker(
|
|
bind=engine,
|
|
expire_on_commit=False,
|
|
autoflush=False,
|
|
autocommit=False,
|
|
)
|
|
|
|
|
|
@when("I save the trace inside a UnitOfWork transaction that rolls back")
|
|
def step_save_in_uow_rollback(context: Context) -> None:
|
|
session = context.uow_session_factory()
|
|
repo = LLMTraceRepository(session_factory=lambda: session)
|
|
try:
|
|
# Pass the session explicitly to use UoW mode: save() flushes but
|
|
# does NOT commit, so the caller's rollback can undo the change.
|
|
repo.save(context.trace, session=session)
|
|
# Simulate a subsequent failure that triggers rollback
|
|
raise RuntimeError("Simulated failure after save")
|
|
except RuntimeError:
|
|
session.rollback()
|
|
finally:
|
|
session.close()
|
|
|
|
|
|
@then("the LLM trace should not be persisted in the database")
|
|
def step_trace_not_persisted(context: Context) -> None:
|
|
session = context.uow_session_factory()
|
|
try:
|
|
row = (
|
|
session.query(LLMTraceModel)
|
|
.filter_by(trace_id=context.trace.trace_id)
|
|
.first()
|
|
)
|
|
assert row is None, (
|
|
f"Expected LLM trace {context.trace.trace_id} to be rolled back, "
|
|
f"but it was found in the database"
|
|
)
|
|
finally:
|
|
session.close()
|