diff --git a/CHANGELOG.md b/CHANGELOG.md index d8335195..8301c6f1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,11 +7,19 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- Retriever logging accepts an explicit `data_source_id`, and agent logging + accepts `span_kind=SpanKind.CLIENT` for calls to remote agent services. + ### Fixed - Agent Control spans exported over OTLP now include the control discriminator and complete `agent_control.*` field set required for backend classification and Controls-card rendering. +- Retriever spans exported over OTLP now use client operation semantics and are + named `retrieval {data_source_id}` when an explicit data-source ID is supplied, + falling back to the captured display name and then to `retrieval`. ## [0.1.1] - 2026-08-03 diff --git a/src/splunk_ao/converter/attribute_mapping.py b/src/splunk_ao/converter/attribute_mapping.py index f753f818..88842f75 100644 --- a/src/splunk_ao/converter/attribute_mapping.py +++ b/src/splunk_ao/converter/attribute_mapping.py @@ -423,6 +423,7 @@ def set_retriever_attributes(attrs: MutableMapping[str, AttributeValue], span: R attrs["gen_ai.retrieval.documents"] = _content_value(span.output) attrs["splunk_ao.retrieval.documents.count"] = len(span.output) attrs["db.operation"] = "search" + _set_if_present(attrs, "gen_ai.data_source.id", _field(span, "data_source_id")) requested_top_k = _field(span, "num_documents") _set_if_present(attrs, "gen_ai.retrieval.top_k", requested_top_k) diff --git a/src/splunk_ao/converter/span_converter.py b/src/splunk_ao/converter/span_converter.py index d041a482..fef61f6d 100644 --- a/src/splunk_ao/converter/span_converter.py +++ b/src/splunk_ao/converter/span_converter.py @@ -20,16 +20,12 @@ _NAME_PARTS: dict[StepType, tuple[str, str]] = { StepType.llm: ("chat", "model"), StepType.tool: ("execute_tool", "name"), - StepType.retriever: ("retrieval", "name"), + StepType.retriever: ("retrieval", "data_source_id"), StepType.workflow: ("invoke_workflow", "name"), StepType.agent: ("invoke_agent", "name"), StepType.control: ("", "name"), } -_KIND_BY_STEP_TYPE = { - step_type: SpanKind.CLIENT if step_type is StepType.llm else SpanKind.INTERNAL for step_type in _NAME_PARTS -} - def _step_type(span: BaseStep) -> StepType: raw_type = getattr(span, "type", None) @@ -48,9 +44,21 @@ def _step_type(span: BaseStep) -> StepType: def _span_name(span: BaseStep, step_type: StepType) -> str: prefix, field_name = _NAME_PARTS[step_type] detail = getattr(span, field_name, None) + if step_type is StepType.retriever: + if detail is not None: + return f"{prefix} {detail}" + detail = getattr(span, "name", None) return " ".join(part for part in (prefix, str(detail).strip() if detail is not None else "") if part) +def _span_kind(span: BaseStep, step_type: StepType) -> SpanKind: + if step_type in (StepType.llm, StepType.retriever): + return SpanKind.CLIENT + if step_type is StepType.agent and getattr(span, "span_kind", None) is SpanKind.CLIENT: + return SpanKind.CLIENT + return SpanKind.INTERNAL + + def _to_unix_ns(value: datetime) -> int: if value.tzinfo is None: value = value.replace(tzinfo=UTC) @@ -98,7 +106,7 @@ def convert_span( attributes=build_span_attributes(span, session_id), events=(), links=(), - kind=_KIND_BY_STEP_TYPE[step_type], + kind=_span_kind(span, step_type), instrumentation_scope=_INSTRUMENTATION_SCOPE, status=_span_status(span), start_time=start_time_ns, diff --git a/src/splunk_ao/decorator.py b/src/splunk_ao/decorator.py index 99dd26ec..738c87d2 100644 --- a/src/splunk_ao/decorator.py +++ b/src/splunk_ao/decorator.py @@ -55,6 +55,7 @@ def call_llm(prompt, temperature=0.7): from types import TracebackType from typing import Any, TypeVar, cast, overload +from opentelemetry.trace import SpanKind from typing_extensions import ParamSpec from galileo_core.schemas.logging.span import WorkflowSpan @@ -276,6 +277,8 @@ def log( name: str | None = None, span_type: SPAN_TYPE | None = None, params: dict[str, str | Callable] | None = None, + data_source_id: str | None = None, + span_kind: SpanKind | None = None, ) -> Callable[[Callable[P, R]], Callable[P, R]]: ... def log( @@ -286,6 +289,8 @@ def log( span_type: SPAN_TYPE | None = None, params: dict[str, str | Callable] | None = None, dataset_record: DatasetRecord | None = None, + data_source_id: str | None = None, + span_kind: SpanKind | None = None, ) -> Callable[[Callable[P, R]], Callable[P, R]]: """ Main decorator function for logging function calls. @@ -304,6 +309,10 @@ def log( Optional span type ("llm", "retriever", "tool", "workflow", "agent") params Optional parameter mapping for extracting specific values + data_source_id + Optional authoritative retrieval data-source ID for ``span_type="retriever"``. + span_kind + Optional OTel kind for ``span_type="agent"``. Only ``SpanKind.CLIENT`` marks a remote agent call. dataset_record Optional parameter for dataset values. This is used by the local experiment module to set the dataset fields on the trace/spans and not generally provided for logging to log streams. @@ -311,26 +320,55 @@ def log( ------- A decorated function that logs its execution """ + explicit_span_params = { + key: value + for key, value in {"data_source_id": data_source_id, "span_kind": span_kind}.items() + if value is not None + } def decorator(func: Callable[P, R]) -> Callable[P, R]: if inspect.isasyncgenfunction(func): return cast( Callable[P, R], self._async_generator_log( - func, name=name, span_type=span_type, params=params, dataset_record=dataset_record + func, + name=name, + span_type=span_type, + params=params, + dataset_record=dataset_record, + explicit_span_params=explicit_span_params, ), ) if inspect.isgeneratorfunction(func): return cast( Callable[P, R], self._sync_generator_log( - func, name=name, span_type=span_type, params=params, dataset_record=dataset_record + func, + name=name, + span_type=span_type, + params=params, + dataset_record=dataset_record, + explicit_span_params=explicit_span_params, ), ) wrapped = ( - self._async_log(func, name=name, span_type=span_type, params=params, dataset_record=dataset_record) + self._async_log( + func, + name=name, + span_type=span_type, + params=params, + dataset_record=dataset_record, + explicit_span_params=explicit_span_params, + ) if asyncio.iscoroutinefunction(func) - else self._sync_log(func, name=name, span_type=span_type, params=params, dataset_record=dataset_record) + else self._sync_log( + func, + name=name, + span_type=span_type, + params=params, + dataset_record=dataset_record, + explicit_span_params=explicit_span_params, + ) ) return cast(Callable[P, R], wrapped) @@ -348,6 +386,7 @@ def _async_log( span_type: SPAN_TYPE | None, params: dict[str, str | Callable] | None = None, dataset_record: DatasetRecord | None = None, + explicit_span_params: dict[str, Any] | None = None, ) -> F: """ Internal method to handle logging for async functions. @@ -384,6 +423,7 @@ async def async_wrapper(*args: Any, **kwargs: Any) -> Any: name=name or func.__name__, span_type=span_type, params=params, + explicit_span_params=explicit_span_params, is_method=self._is_method(func), func_args=args, func_kwargs=kwargs, @@ -415,6 +455,7 @@ def _sync_log( span_type: SPAN_TYPE | None, params: dict[str, str | Callable] | None = None, dataset_record: DatasetRecord | None = None, + explicit_span_params: dict[str, Any] | None = None, ) -> F: """ Internal method to handle logging for synchronous functions. @@ -442,6 +483,7 @@ def sync_wrapper(*args: Any, **kwargs: Any) -> Any: name=name or func.__name__, span_type=span_type, params=params, + explicit_span_params=explicit_span_params, is_method=self._is_method(func), func_args=args, func_kwargs=kwargs, @@ -473,6 +515,7 @@ def _sync_generator_log( span_type: SPAN_TYPE | None, params: dict[str, str | Callable] | None = None, dataset_record: DatasetRecord | None = None, + explicit_span_params: dict[str, Any] | None = None, ) -> F: @wraps(func) def generator_wrapper(*args: Any, **kwargs: Any) -> Generator: @@ -481,6 +524,7 @@ def generator_wrapper(*args: Any, **kwargs: Any) -> Generator: name=name or func.__name__, span_type=span_type, params=params, + explicit_span_params=explicit_span_params, is_method=self._is_method(func), func_args=args, func_kwargs=kwargs, @@ -500,6 +544,7 @@ def _async_generator_log( span_type: SPAN_TYPE | None, params: dict[str, str | Callable] | None = None, dataset_record: DatasetRecord | None = None, + explicit_span_params: dict[str, Any] | None = None, ) -> F: @wraps(func) def async_generator_wrapper(*args: Any, **kwargs: Any) -> AsyncGenerator: @@ -508,6 +553,7 @@ def async_generator_wrapper(*args: Any, **kwargs: Any) -> AsyncGenerator: name=name or func.__name__, span_type=span_type, params=params, + explicit_span_params=explicit_span_params, is_method=self._is_method(func), func_args=args, func_kwargs=kwargs, @@ -543,6 +589,7 @@ def _prepare_input( name: str, span_type: SPAN_TYPE | None, params: dict[str, str | Callable] | None = None, + explicit_span_params: dict[str, Any] | None = None, is_method: bool = False, func_args: tuple = (), func_kwargs: dict | None = None, @@ -606,6 +653,8 @@ def _prepare_input( if param_name in input_ and param_name not in span_params: span_params[param_name] = input_[param_name] + span_params.update(explicit_span_params or {}) + if "name" not in span_params: span_params["name"] = name @@ -685,10 +734,10 @@ def _get_span_param_names(self, span_type: SPAN_TYPE) -> list[str]: common_params = ["name", "input", "metadata", "tags"] span_params = { "llm": [*common_params, "model", "temperature", "tools"], - "retriever": common_params, + "retriever": [*common_params, "data_source_id"], "tool": [*common_params, "tool_call_id"], "workflow": common_params, - "agent": [*common_params, "agent_type"], + "agent": [*common_params, "agent_type", "span_kind"], } return span_params.get(span_type, common_params) @@ -791,8 +840,9 @@ def _prepare_call( created_at = span_params.get("created_at", _get_timestamp()) if span_type == "agent": agent_type = span_params.get("agent_type") + span_kind = span_params.get("span_kind", SpanKind.INTERNAL) span = client_instance.add_agent_span( - input=input_, name=name, agent_type=agent_type, created_at=created_at + input=input_, name=name, agent_type=agent_type, span_kind=span_kind, created_at=created_at ) else: span = client_instance.add_workflow_span(input=input_, name=name, created_at=created_at) diff --git a/src/splunk_ao/logger/logger.py b/src/splunk_ao/logger/logger.py index 8ebd8937..c7c5bb2c 100644 --- a/src/splunk_ao/logger/logger.py +++ b/src/splunk_ao/logger/logger.py @@ -20,7 +20,7 @@ from opentelemetry import context as otel_context from opentelemetry import trace as otel_trace from opentelemetry.sdk.trace.id_generator import RandomIdGenerator -from opentelemetry.trace import NonRecordingSpan, SpanContext, TraceFlags, TraceState +from opentelemetry.trace import NonRecordingSpan, SpanContext, SpanKind, TraceFlags, TraceState from pydantic import PrivateAttr from galileo_core.helpers.execution import async_run @@ -31,7 +31,6 @@ LlmSpan, LlmSpanAllowedInputType, LlmSpanAllowedOutputType, - RetrieverSpan, Span, StepWithChildSpans, ToolSpan, @@ -69,6 +68,7 @@ LoggedAgentSpan, LoggedControlSpan, LoggedLlmSpan, + LoggedRetrieverSpan, LoggedTrace, LoggedWorkflowSpan, TextOrContentBlocks, @@ -368,7 +368,9 @@ def __init__( "User must provide project_name or project_id to SplunkAOLogger, or set it as an environment variable." ) if self.experiment_id is None and self.agent_stream_name is None and self.agent_stream_id is None: - raise SplunkAOLoggerException("agent_stream or agent_stream_id is required to initialize SplunkAOLogger.") + raise SplunkAOLoggerException( + "agent_stream or agent_stream_id is required to initialize SplunkAOLogger." + ) if local_metrics: self.local_metrics = local_metrics @@ -1705,7 +1707,8 @@ def add_retriever_span( tags: list[str] | None = None, status_code: int | None = None, step_number: int | None = None, - ) -> RetrieverSpan: + data_source_id: str | None = None, + ) -> LoggedRetrieverSpan: """ Add a new retriever span to the current parent. @@ -1738,10 +1741,12 @@ def add_retriever_span( Status code of the node execution. step_number: Optional[int] Step number of the span. + data_source_id: Optional[str] + Authoritative identifier of the retrieval data source used by the GenAI system. Returns ------- - RetrieverSpan + LoggedRetrieverSpan The created span. """ documents = convert_to_documents(output, "output") @@ -1751,22 +1756,23 @@ def add_retriever_span( if metadata: metadata = {k: SplunkAOLogger._convert_metadata_value(v) for k, v in metadata.items()} - kwargs = { - "input": input, - "documents": documents, - "redacted_input": redacted_input, - "redacted_documents": redacted_documents, - "name": name, - "duration_ns": duration_ns, - "created_at": created_at, - "user_metadata": metadata, - "tags": tags, - "status_code": status_code, - "step_number": step_number, - "id": uuid.uuid4(), - } parent = self.current_parent() - span = super().add_retriever_span(**kwargs) + span = LoggedRetrieverSpan( + input=input, + output=documents, + redacted_input=redacted_input, + redacted_output=redacted_documents, + name=name, + metrics=Metrics(duration_ns=duration_ns), + created_at=self._get_child_span_timestamp() if created_at is None else created_at, + user_metadata=metadata, + tags=tags, + status_code=status_code, + step_number=step_number, + data_source_id=data_source_id, + id=uuid.uuid4(), + ) + self.add_child_span_to_parent(span) span._parent = parent self._record_otel_ids(span, parent_step=parent) self._mark_potentially_parentable(span) @@ -1969,6 +1975,7 @@ def add_agent_span( agent_type: AgentType | None = None, step_number: int | None = None, status_code: int | None = None, + span_kind: SpanKind = SpanKind.INTERNAL, ) -> LoggedAgentSpan: """ Add an agent type span to the current parent. @@ -2010,6 +2017,9 @@ def add_agent_span( Step number of the span. status_code: Optional[int] Status code of the span execution (e.g., 200 for success, 500 for error). + span_kind: SpanKind + OTel operation kind. Only ``SpanKind.CLIENT`` is accepted as a remote-agent override; + all other values use ``SpanKind.INTERNAL``. Returns ------- @@ -2031,6 +2041,7 @@ def add_agent_span( tags=tags, metrics=Metrics(duration_ns=duration_ns), agent_type=agent_type, + span_kind=span_kind, id=uuid.uuid4(), step_number=step_number, ) diff --git a/src/splunk_ao/schema/logged.py b/src/splunk_ao/schema/logged.py index 01565318..9ee528fb 100644 --- a/src/splunk_ao/schema/logged.py +++ b/src/splunk_ao/schema/logged.py @@ -9,7 +9,8 @@ from json import dumps from typing import Annotated, Any -from pydantic import Field +from opentelemetry.trace import SpanKind +from pydantic import ConfigDict, Field, field_validator from galileo_core.schemas.logging.llm import Message, MessageRole from galileo_core.schemas.logging.span import ( @@ -73,6 +74,32 @@ class LoggedAgentSpan(AgentSpan): redacted_output: IngestOutputType | None = _REDACTED_OUTPUT_FIELD spans: list["LoggedSpan"] = Field(default_factory=list) conversation_root: bool | None = Field(default=None) + span_kind: SpanKind = Field(default=SpanKind.INTERNAL, exclude=True) + + @field_validator("span_kind", mode="before") + @classmethod + def normalize_span_kind(cls, value: object) -> SpanKind: + """Allow only the explicit remote-agent client classification.""" + return SpanKind.CLIENT if value is SpanKind.CLIENT else SpanKind.INTERNAL + + +class LoggedRetrieverSpan(RetrieverSpan): + """RetrieverSpan with SDK-local OTel data-source identity.""" + + # LoggedTrace accepts existing core RetrieverSpan instances and widens them + # through discriminated-union validation. + model_config = ConfigDict(from_attributes=True) + + data_source_id: str | None = Field(default=None, exclude=True) + + @field_validator("data_source_id", mode="before") + @classmethod + def normalize_data_source_id(cls, value: object) -> object: + """Trim a data-source ID and treat a blank value as absent.""" + if isinstance(value, str): + value = value.strip() + return value or None + return value class LoggedLlmSpan(LlmSpan): @@ -128,15 +155,16 @@ class LoggedControlSpan(ControlSpan): redacted_input: str | None = Field(default=None, description=BaseStep.model_fields["redacted_input"].description) -# RetrieverSpan and ToolSpan use plain string/document I/O and don't need multimodal widening. +# ToolSpan uses plain string I/O and doesn't need multimodal widening. # LoggedControlSpan narrows ControlSpan's Any input to str for the same reason. LoggedSpan = Annotated[ - LoggedAgentSpan | LoggedWorkflowSpan | LoggedLlmSpan | RetrieverSpan | ToolSpan | LoggedControlSpan, + LoggedAgentSpan | LoggedWorkflowSpan | LoggedLlmSpan | LoggedRetrieverSpan | ToolSpan | LoggedControlSpan, Field(discriminator="type"), ] LoggedTrace.model_rebuild() LoggedWorkflowSpan.model_rebuild() LoggedAgentSpan.model_rebuild() +LoggedRetrieverSpan.model_rebuild() LoggedLlmSpan.model_rebuild() LoggedControlSpan.model_rebuild() diff --git a/tests/schemas/test_logged.py b/tests/schemas/test_logged.py index 2d1a24c0..0aeb146d 100644 --- a/tests/schemas/test_logged.py +++ b/tests/schemas/test_logged.py @@ -1,6 +1,7 @@ """Tests for SDK-local ingestion models (Logged variants and content blocks).""" import pytest +from opentelemetry.trace import SpanKind from pydantic import ValidationError from galileo_core.schemas.logging.llm import MessageRole @@ -9,7 +10,14 @@ from galileo_core.schemas.shared.document import Document from galileo_core.schemas.shared.multimodal import ContentModality from splunk_ao.schema.content_blocks import DataContentBlock, TextContentBlock -from splunk_ao.schema.logged import LoggedAgentSpan, LoggedControlSpan, LoggedLlmSpan, LoggedTrace, LoggedWorkflowSpan +from splunk_ao.schema.logged import ( + LoggedAgentSpan, + LoggedControlSpan, + LoggedLlmSpan, + LoggedRetrieverSpan, + LoggedTrace, + LoggedWorkflowSpan, +) from splunk_ao.schema.message import LoggedMessage from splunk_ao.schema.trace import TracesIngestRequest @@ -238,13 +246,14 @@ def test_retriever_span_roundtrip(self) -> None: ) ], ) + assert type(trace.spans[0]) is LoggedRetrieverSpan # When: JSON roundtrip raw = trace.model_dump(mode="json") restored = LoggedTrace.model_validate(raw) # Then: RetrieverSpan and its Document output are preserved - assert type(restored.spans[0]) is RetrieverSpan + assert type(restored.spans[0]) is LoggedRetrieverSpan assert restored.spans[0].output[0].content == "retrieved doc" def test_tool_span_roundtrip(self) -> None: @@ -332,7 +341,7 @@ def test_full_ingest_request_roundtrip(self) -> None: assert llm.input[0].content[1].base64 == "audio_b64" retriever = wf.spans[1] - assert type(retriever) is RetrieverSpan + assert type(retriever) is LoggedRetrieverSpan assert retriever.output[0].content == "ctx doc" tool = wf.spans[2] @@ -341,6 +350,7 @@ def test_full_ingest_request_roundtrip(self) -> None: def test_logged_trace_roundtrip_with_control_span(self) -> None: from splunk_ao.logger.control import ControlSpan + # Given: a native Core ControlSpan payload control_payload = ControlSpan(input="selected text").model_dump(mode="python") trace = LoggedTrace(input="query", spans=[control_payload]) @@ -354,6 +364,7 @@ def test_logged_trace_roundtrip_with_control_span(self) -> None: @pytest.mark.parametrize("field_name", ["id", "session_id", "trace_id", "parent_id"]) def test_control_span_rejects_non_uuidish_id_fields(self, field_name: str) -> None: from splunk_ao.logger.control import ControlSpan + # When/Then: non-UUID-ish ID fields are rejected by the native Core schema with pytest.raises(ValidationError): ControlSpan(input="selected text", **{field_name: 123}) @@ -419,6 +430,24 @@ def test_logged_spans_are_instances_of_core(self) -> None: assert isinstance(LoggedWorkflowSpan(input="x"), WorkflowSpan) assert isinstance(LoggedAgentSpan(input="x"), AgentSpan) assert isinstance(LoggedLlmSpan(input="x"), LlmSpan) + assert isinstance(LoggedRetrieverSpan(input="x"), RetrieverSpan) + + def test_otel_runtime_hints_are_excluded_from_legacy_ingestion_payload(self) -> None: + # Given: path-1 spans carrying the new SDK-local OTel hints + trace = LoggedTrace( + input="input", + spans=[ + LoggedAgentSpan(input="input", span_kind=SpanKind.CLIENT), + LoggedRetrieverSpan(input="query", data_source_id="knowledge-base"), + ], + ) + + # When: the deprecated proprietary ingestion request is serialized + spans = TracesIngestRequest(traces=[trace]).model_dump(mode="json")["traces"][0]["spans"] + + # Then: neither runtime-only hint expands that wire contract + assert "span_kind" not in spans[0] + assert "data_source_id" not in spans[1] class TestLoggedTraceTypeRestrictions: diff --git a/tests/test_attribute_mapping.py b/tests/test_attribute_mapping.py index 5fff65ec..4dbe8cae 100644 --- a/tests/test_attribute_mapping.py +++ b/tests/test_attribute_mapping.py @@ -26,6 +26,7 @@ ) from splunk_ao.logger.control import ControlAppliesTo, ControlCheckStage, ControlResult, ControlSpan from splunk_ao.schema import DataContentBlock, LoggedControlSpan, LoggedLlmSpan, LoggedMessage, TextContentBlock +from splunk_ao.schema.logged import LoggedRetrieverSpan def _text_message(role: str, content: str, *, finish_reason: str | None = None) -> dict: @@ -234,9 +235,20 @@ def test_retriever_mapping_uses_query_and_documents() -> None: assert json.loads(attrs["gen_ai.retrieval.documents"]) == [{"content": "doc", "metadata": {"source": "kb"}}] assert attrs["splunk_ao.retrieval.documents.count"] == 1 assert attrs["db.operation"] == "search" + assert "gen_ai.data_source.id" not in attrs assert "gen_ai.output.messages" not in attrs +def test_retriever_mapping_uses_only_explicit_data_source_id() -> None: + span = LoggedRetrieverSpan( + name="display-name", data_source_id="vector-db", input="what is RAG?", output=[Document(content="doc")] + ) + + attrs = build_span_attributes(span) + + assert attrs["gen_ai.data_source.id"] == "vector-db" + + @pytest.mark.parametrize( ("span", "operation_key", "operation", "name_key", "name"), [ @@ -558,11 +570,7 @@ def test_control_mapping_omits_unpopulated_optional_fields() -> None: def test_control_mapping_accepts_schema_compatible_control_span() -> None: - source = ControlSpan( - name="guardrail", - output=ControlResult(action="observe", matched=True), - control_id=42, - ) + source = ControlSpan(name="guardrail", output=ControlResult(action="observe", matched=True), control_id=42) alternate = SimpleNamespace( **{field_name: getattr(source, field_name) for field_name in type(source).model_fields}, model_extra={} ) @@ -601,8 +609,7 @@ def test_control_mapping_tolerates_span_without_control_fields() -> None: def test_control_mapping_exports_error_result_without_dropping_false() -> None: span = ControlSpan( - name="guardrail", - output=ControlResult(action="observe", matched=False, error_message="evaluator unavailable"), + name="guardrail", output=ControlResult(action="observe", matched=False, error_message="evaluator unavailable") ) attrs = build_span_attributes(span) diff --git a/tests/test_decorator.py b/tests/test_decorator.py index f95167d3..b23e3097 100644 --- a/tests/test_decorator.py +++ b/tests/test_decorator.py @@ -3,6 +3,7 @@ from uuid import UUID import pytest +from opentelemetry.trace import SpanKind from pydantic import BaseModel from galileo_core.schemas.logging.span import AgentSpan, LlmSpan, RetrieverSpan, ToolSpan, WorkflowSpan @@ -11,6 +12,7 @@ from splunk_ao import Message, MessageRole, log, splunk_ao_context, start_session from splunk_ao.decorator import _session_id_context from splunk_ao.schema.content_blocks import DataContentBlock, TextContentBlock +from splunk_ao.schema.logged import LoggedAgentSpan, LoggedRetrieverSpan from tests.testutils.setup import setup_mock_logstreams_client, setup_mock_projects_client, setup_mock_traces_client @@ -521,6 +523,54 @@ def retriever_call(query: str) -> str: assert payload.traces[0].spans[0].output == [Document(content="response1", metadata=None)] +@patch("splunk_ao.logger.logger.AgentStreams") +@patch("splunk_ao.logger.logger.Projects") +@patch("splunk_ao.logger.logger.Traces") +def test_decorator_retriever_accepts_explicit_data_source_id( + mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock, reset_context +) -> None: + # Given: a retriever decorator with an authoritative data-source ID + setup_mock_traces_client(mock_traces_client) + setup_mock_projects_client(mock_projects_client) + setup_mock_logstreams_client(mock_logstreams_client) + + @log(span_type="retriever", data_source_id="knowledge-base") + def retriever_call(query: str) -> str: + return "result" + + # When: the decorated retrieval runs + retriever_call(query="input") + + # Then: the runtime-only ID reaches the path-1 retriever model + span = splunk_ao_context.get_logger_instance().traces[-1].spans[0] + assert type(span) is LoggedRetrieverSpan + assert span.data_source_id == "knowledge-base" + + +@patch("splunk_ao.logger.logger.AgentStreams") +@patch("splunk_ao.logger.logger.Projects") +@patch("splunk_ao.logger.logger.Traces") +def test_decorator_agent_accepts_explicit_client_kind( + mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock, reset_context +) -> None: + # Given: an agent decorator explicitly representing a remote call + setup_mock_traces_client(mock_traces_client) + setup_mock_projects_client(mock_projects_client) + setup_mock_logstreams_client(mock_logstreams_client) + + @log(span_type="agent", span_kind=SpanKind.CLIENT) + def agent_call(query: str) -> str: + return "result" + + # When: the decorated agent runs + agent_call(query="input") + + # Then: the runtime-only client kind reaches the path-1 agent model + span = splunk_ao_context.get_logger_instance().traces[-1].spans[0] + assert type(span) is LoggedAgentSpan + assert span.span_kind is SpanKind.CLIENT + + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") diff --git a/tests/test_logger_otel_egress.py b/tests/test_logger_otel_egress.py index 2a3a976c..2672a321 100644 --- a/tests/test_logger_otel_egress.py +++ b/tests/test_logger_otel_egress.py @@ -6,6 +6,7 @@ import pytest from opentelemetry import context, propagate, trace from opentelemetry.sdk.trace import ReadableSpan +from opentelemetry.trace import SpanKind from splunk_ao.deployment import DeploymentMode from splunk_ao.exceptions import SplunkAOLoggerException @@ -108,6 +109,45 @@ def test_pending_leaf_emits_when_trace_envelope_is_released( assert otlp_logger._otel_ids == {} +def test_retriever_emits_client_semantics_from_explicit_data_source_id( + otlp_logger: SplunkAOLogger, recording_sink: RecordingSink +) -> None: + # Given: an active trace and an explicit authoritative data-source ID + otlp_logger.start_trace(input="question") + + # When: a retriever is added and the trace is concluded + otlp_logger.add_retriever_span( + input="query", output=["document"], name="display-name", data_source_id="knowledge-base" + ) + otlp_logger.conclude(output="answer") + + # Then: path 1 emits canonical retrieval name, kind, and identity + [span] = recording_sink.spans + assert span.name == "retrieval knowledge-base" + assert span.kind is SpanKind.CLIENT + assert (span.attributes or {})["gen_ai.data_source.id"] == "knowledge-base" + + +@pytest.mark.parametrize( + ("requested_kind", "expected_kind"), [(SpanKind.CLIENT, SpanKind.CLIENT), (SpanKind.SERVER, SpanKind.INTERNAL)] +) +def test_agent_logger_allows_only_client_kind_override( + otlp_logger: SplunkAOLogger, recording_sink: RecordingSink, requested_kind: SpanKind, expected_kind: SpanKind +) -> None: + # Given: an active trace + otlp_logger.start_trace(input="question") + + # When: an agent is logged with the requested OTel kind + otlp_logger.add_agent_span(input="question", name="agent", span_kind=requested_kind) + # Conclude the agent span, then the trace envelope. + otlp_logger.conclude(output="answer") + otlp_logger.conclude(output="answer") + + # Then: only CLIENT survives as a remote-agent classification + [span] = recording_sink.spans + assert span.kind is expected_kind + + def test_promoted_tool_emits_after_child_and_before_root( otlp_logger: SplunkAOLogger, recording_sink: RecordingSink ) -> None: diff --git a/tests/test_span_converter.py b/tests/test_span_converter.py index 61d93b06..bb6bbdbc 100644 --- a/tests/test_span_converter.py +++ b/tests/test_span_converter.py @@ -15,7 +15,7 @@ from splunk_ao.converter import SpanConverter, span_converter from splunk_ao.converter.attribute_mapping import build_span_attributes from splunk_ao.logger.control import ControlAppliesTo, ControlCheckStage, ControlResult, ControlSpan -from splunk_ao.schema.logged import LoggedControlSpan +from splunk_ao.schema.logged import LoggedAgentSpan, LoggedControlSpan, LoggedRetrieverSpan from splunk_ao.utils.headers_data import get_package_version TRACE_ID = 0x1234567890ABCDEF1234567890ABCDEF @@ -73,7 +73,7 @@ def supported_spans() -> list[tuple[BaseStep, str, SpanKind]]: name="knowledge-base", input="query", output=[Document(content="result")], created_at=CREATED_AT ), "retrieval knowledge-base", - SpanKind.INTERNAL, + SpanKind.CLIENT, ), ( WorkflowSpan(name="research", input="question", output="answer", created_at=CREATED_AT), @@ -124,6 +124,67 @@ def test_missing_optional_name_parts_do_not_leave_whitespace(span: BaseStep, exp assert convert(span).name == expected_name +def test_retriever_name_and_attribute_use_explicit_data_source_id_byte_for_byte() -> None: + span = LoggedRetrieverSpan( + name="display-name", data_source_id="knowledge base/v1", input="query", output=[Document(content="result")] + ) + + result = convert(span) + + assert result.name == "retrieval knowledge base/v1" + assert result.kind is SpanKind.CLIENT + assert (result.attributes or {})["gen_ai.data_source.id"] == "knowledge base/v1" + + +def test_retriever_display_name_is_used_only_as_span_name_fallback() -> None: + span = LoggedRetrieverSpan(name="display-name", input="query", output=[]) + + result = convert(span) + + assert result.name == "retrieval display-name" + assert "gen_ai.data_source.id" not in (result.attributes or {}) + + +@pytest.mark.parametrize("data_source_id", ["", " ", "\t\n"]) +def test_blank_retriever_data_source_id_is_treated_as_absent(data_source_id: str) -> None: + # Given: an explicitly blank data-source ID + span = LoggedRetrieverSpan(name="display-name", data_source_id=data_source_id, input="query", output=[]) + + # When: the proprietary retriever is converted + result = convert(span) + + # Then: the display name remains visible without emitting a blank semantic attribute + assert result.name == "retrieval display-name" + assert "gen_ai.data_source.id" not in (result.attributes or {}) + + +def test_padded_retriever_data_source_id_is_normalized_consistently() -> None: + # Given: a data-source ID padded with accidental whitespace + span = LoggedRetrieverSpan(name="display-name", data_source_id=" knowledge base/v1 ", input="query", output=[]) + + # When: the proprietary retriever is converted + result = convert(span) + + # Then: the normalized ID is used consistently in the name and attribute + assert result.name == "retrieval knowledge base/v1" + assert (result.attributes or {})["gen_ai.data_source.id"] == "knowledge base/v1" + + +@pytest.mark.parametrize( + ("requested_kind", "expected_kind"), + [ + (SpanKind.INTERNAL, SpanKind.INTERNAL), + (SpanKind.CLIENT, SpanKind.CLIENT), + (SpanKind.SERVER, SpanKind.INTERNAL), + ("CLIENT", SpanKind.INTERNAL), + ], +) +def test_agent_kind_allows_only_explicit_client_override(requested_kind: object, expected_kind: SpanKind) -> None: + span = LoggedAgentSpan(name="agent", span_kind=requested_kind) + + assert convert(span).kind is expected_kind + + @pytest.mark.parametrize( "span", [BaseStep(type=StepType.session), ToolSpan(name="tool").model_copy(update={"type": "unknown"})] )