From 1f9eddf968e6e87edd46b6c3f4f5391350e2d301 Mon Sep 17 00:00:00 2001 From: adityamehra Date: Fri, 17 Jul 2026 16:14:56 -0700 Subject: [PATCH 1/4] =?UTF-8?q?feat(HYBIM-730):=20rename=20Log=20Streams?= =?UTF-8?q?=20=E2=86=92=20Agent=20Streams,=20Metrics=20=E2=86=92=20Evaluat?= =?UTF-8?q?ors?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Introduces the new canonical domain entity names required by HYBIM-730: - **AgentStream** replaces LogStream (new class in agent_stream.py that subclasses LogStream; returns AgentStream from Project methods) - **AgentStreams** replaces LogStreams service class (agent_streams.py); module-level helpers: get_agent_stream, list_agent_streams, create_agent_stream, enable_evaluators - **Evaluator** replaces Metric base class (evaluator.py) - **LlmEvaluator**, **CodeEvaluator**, **LocalEvaluator**, **SplunkAOEvaluator** replace their Metric-prefixed counterparts - **Evaluators** replaces Metrics service class (evaluators.py); helpers: create_custom_llm_evaluator, get_evaluators, delete_evaluator Backward compatibility: - Old names (LogStream, Metric, LlmMetric, …) are preserved via PEP 562 __getattr__ in __init__.py and emit DeprecationWarning - Project.create_log_stream / list_log_streams / logstreams delegate to the new methods with DeprecationWarning - __future__/agent_stream.py and __future__/evaluator.py shims added The underlying API endpoints (/log_streams, /scorers) are unchanged; the server-side rename is tracked separately. Closes HYBIM-730 Co-authored-by: Cursor --- docs/HYBIM-730-domain-entity-rename.md | 237 +++++++++++++++++++++++ src/splunk_ao/__future__/agent_stream.py | 14 ++ src/splunk_ao/__future__/evaluator.py | 20 ++ src/splunk_ao/__init__.py | 63 +++++- src/splunk_ao/agent_stream.py | 61 ++++++ src/splunk_ao/agent_streams.py | 208 ++++++++++++++++++++ src/splunk_ao/evaluator.py | 188 ++++++++++++++++++ src/splunk_ao/evaluators.py | 194 +++++++++++++++++++ src/splunk_ao/project.py | 137 ++++++++++--- 9 files changed, 1082 insertions(+), 40 deletions(-) create mode 100644 docs/HYBIM-730-domain-entity-rename.md create mode 100644 src/splunk_ao/__future__/agent_stream.py create mode 100644 src/splunk_ao/__future__/evaluator.py create mode 100644 src/splunk_ao/agent_stream.py create mode 100644 src/splunk_ao/agent_streams.py create mode 100644 src/splunk_ao/evaluator.py create mode 100644 src/splunk_ao/evaluators.py diff --git a/docs/HYBIM-730-domain-entity-rename.md b/docs/HYBIM-730-domain-entity-rename.md new file mode 100644 index 00000000..5afcf215 --- /dev/null +++ b/docs/HYBIM-730-domain-entity-rename.md @@ -0,0 +1,237 @@ +# Domain Entity Rename: Log Streams → Agent Streams, Metrics → Evaluators + +**Ticket:** [HYBIM-730](https://splunk.atlassian.net/browse/HYBIM-730) +**Branch:** `feat/HYBIM-730-domain-rename` +**Status:** Implemented with full backward compatibility + +--- + +## Summary + +Two core SDK domain entities are being renamed to better reflect their purpose: + +| Old name | New name | Scope | +|----------|----------|-------| +| Log Stream / `LogStream` | Agent Stream / `AgentStream` | Top-level class, service class, convenience functions | +| Metric / `Metric` | Evaluator / `Evaluator` | Base class and all concrete subclasses | + +The old names remain functional but emit `DeprecationWarning` at import time. +They will be removed in a future major release. + +> **Note:** The underlying API endpoints continue to use the previous paths +> (`/log_streams`, `/scorers`) — server-side renaming is tracked separately. + +--- + +## New Public API + +### Agent Streams + +```python +# Primary class (replaces LogStream) +from splunk_ao import AgentStream + +# Create and persist a new agent stream +stream = AgentStream(name="prod-traces", project_name="my-project").create() + +# Retrieve an existing stream +stream = AgentStream.get(name="prod-traces", project_name="my-project") + +# List streams for a project +streams = AgentStream.list(project_name="my-project") + +# Enable evaluators on the stream +from splunk_ao import SplunkAOMetrics +stream.set_metrics([SplunkAOMetrics.correctness, SplunkAOMetrics.completeness]) +``` + +```python +# Service class (replaces LogStreams in splunk_ao.log_streams) +from splunk_ao.agent_streams import AgentStreams + +svc = AgentStreams() +stream = svc.get(name="prod-traces", project_name="my-project") +streams = svc.list(project_name="my-project") +``` + +```python +# Convenience functions (replaces functions in splunk_ao.log_streams) +from splunk_ao.agent_streams import ( + get_agent_stream, + list_agent_streams, + create_agent_stream, + enable_evaluators, +) + +stream = get_agent_stream(name="prod-traces", project_name="my-project") +streams = list_agent_streams(project_name="my-project") +stream = create_agent_stream(name="new-stream", project_name="my-project") + +# Enable evaluators (formerly enable_metrics) +local_evals = enable_evaluators( + agent_stream_name="prod-traces", + project_name="my-project", + metrics=[SplunkAOMetrics.correctness, "completeness"], +) +``` + +### Evaluators + +```python +# Base class (replaces Metric) +from splunk_ao import Evaluator, LlmEvaluator, CodeEvaluator, LocalEvaluator, SplunkAOEvaluator + +# Get an existing evaluator by name or ID +ev = Evaluator.get(name="factuality-checker") + +# List all evaluators +evaluators = Evaluator.list() + +# Access built-in evaluators +ev = Evaluator.evaluators.correctness +ev = Evaluator.evaluators.completeness + +# Create a custom LLM evaluator +llm_ev = LlmEvaluator( + name="response_quality", + prompt="Rate the quality 1-10: {input} -> {output}", + model="gpt-4o-mini", + judges=3, +).create() + +# Create a code-based evaluator +code_ev = CodeEvaluator( + name="custom_scorer", + code="def scorer_fn(step): return 1.0", +).create() + +# Create a local function-based evaluator +def my_fn(trace): + return min(len(getattr(trace, "output", "") or "") / 200.0, 1.0) + +local_ev = LocalEvaluator(name="response_length", scorer_fn=my_fn) +``` + +```python +# Evaluators service class (replaces Metrics in splunk_ao.metrics) +from splunk_ao.evaluators import ( + Evaluators, + create_custom_llm_evaluator, + get_evaluators, + delete_evaluator, +) + +create_custom_llm_evaluator(name="my-eval", user_prompt="Rate this...") +delete_evaluator(name="old-eval") +``` + +### Project methods + +```python +from splunk_ao import Project + +project = Project.get(name="My AI Project") + +# New methods +stream = project.create_agent_stream(name="Production Traces") +streams = project.list_agent_streams() +for stream in project.agent_streams: + print(stream.name) +``` + +--- + +## Backward Compatibility + +All old names continue to work but emit `DeprecationWarning`: + +```python +import warnings +warnings.simplefilter("always", DeprecationWarning) + +# These still work — but warn: +from splunk_ao import LogStream # warns: use AgentStream +from splunk_ao import Metric # warns: use Evaluator +from splunk_ao import LlmMetric # warns: use LlmEvaluator +from splunk_ao import CodeMetric # warns: use CodeEvaluator +from splunk_ao import LocalMetric # warns: use LocalEvaluator +from splunk_ao import SplunkAOMetric # warns: use SplunkAOEvaluator + +# Project deprecated methods — warn: +project.create_log_stream("foo") # warns: use create_agent_stream +project.list_log_streams() # warns: use list_agent_streams +project.logstreams # warns: use agent_streams +``` + +### Migration quick-reference + +| Old import | New import | +|------------|------------| +| `from splunk_ao import LogStream` | `from splunk_ao import AgentStream` | +| `from splunk_ao import Metric` | `from splunk_ao import Evaluator` | +| `from splunk_ao import LlmMetric` | `from splunk_ao import LlmEvaluator` | +| `from splunk_ao import CodeMetric` | `from splunk_ao import CodeEvaluator` | +| `from splunk_ao import LocalMetric` | `from splunk_ao import LocalEvaluator` | +| `from splunk_ao import SplunkAOMetric` | `from splunk_ao import SplunkAOEvaluator` | +| `from splunk_ao.log_stream import LogStream` | `from splunk_ao.agent_stream import AgentStream` | +| `from splunk_ao.log_streams import LogStreams` | `from splunk_ao.agent_streams import AgentStreams` | +| `from splunk_ao.log_streams import get_log_stream` | `from splunk_ao.agent_streams import get_agent_stream` | +| `from splunk_ao.log_streams import list_log_streams` | `from splunk_ao.agent_streams import list_agent_streams` | +| `from splunk_ao.log_streams import create_log_stream` | `from splunk_ao.agent_streams import create_agent_stream` | +| `from splunk_ao.log_streams import enable_metrics` | `from splunk_ao.agent_streams import enable_evaluators` | +| `from splunk_ao.metric import Metric` | `from splunk_ao.evaluator import Evaluator` | +| `from splunk_ao.metrics import Metrics` | `from splunk_ao.evaluators import Evaluators` | +| `project.create_log_stream(name)` | `project.create_agent_stream(name)` | +| `project.list_log_streams()` | `project.list_agent_streams()` | +| `project.logstreams` | `project.agent_streams` | + +--- + +## Files Changed + +| File | Change | +|------|--------| +| `src/splunk_ao/agent_stream.py` | **New** — `AgentStream` subclass of `LogStream` | +| `src/splunk_ao/agent_streams.py` | **New** — `AgentStreams` service class + convenience functions | +| `src/splunk_ao/evaluator.py` | **New** — `Evaluator`, `LlmEvaluator`, `CodeEvaluator`, `LocalEvaluator`, `SplunkAOEvaluator` | +| `src/splunk_ao/evaluators.py` | **New** — `Evaluators` service class + convenience functions | +| `src/splunk_ao/__future__/agent_stream.py` | **New** — deprecation shim | +| `src/splunk_ao/__future__/evaluator.py` | **New** — deprecation shim | +| `src/splunk_ao/__init__.py` | **Modified** — export new names; PEP 562 `__getattr__` for deprecated old names | +| `src/splunk_ao/project.py` | **Modified** — add `create_agent_stream`, `list_agent_streams`, `agent_streams`; deprecate old methods | +| `docs/HYBIM-730-domain-entity-rename.md` | **New** — this document | + +--- + +## Design Decisions + +### Subclassing vs type aliases + +`AgentStream` is implemented as a subclass of `LogStream` (rather than a plain +`= LogStream` alias) so that: +- `AgentStream.__name__` is `"AgentStream"` +- Instances created via `AgentStream(...)` have the new name in repr +- Future implementation divergence is possible without API breakage + +The same pattern applies to `Evaluator` / `LlmEvaluator` / etc. + +### PEP 562 `__getattr__` in `__init__.py` + +Old names (`LogStream`, `Metric`, …) are **not** listed in `__all__`, so they +are invisible to tab-completion and static analysis. They are only reachable +via `__getattr__`, which fires a `DeprecationWarning`. This gives existing code +a graceful migration path without polluting the new API surface. + +### Server-side API paths unchanged + +The API endpoints remain on the old paths (`/log_streams`, `/scorers`). +`AgentStream` / `Evaluator` call into the same service layers (`LogStreams`, +`Metrics`) as before. A separate ticket tracks the server-side rename. + +--- + +## Related + +- [HYBIM-730](https://splunk.atlassian.net/browse/HYBIM-730) — this ticket +- [HYBIM-832](https://splunk.atlassian.net/browse/HYBIM-832) — Galileo → Splunk AO env-var rename +- `docs/OTEL_GALILEO_TO_SPLUNK_AO_RENAME.md` — OTel layer rename diff --git a/src/splunk_ao/__future__/agent_stream.py b/src/splunk_ao/__future__/agent_stream.py new file mode 100644 index 00000000..27a7e092 --- /dev/null +++ b/src/splunk_ao/__future__/agent_stream.py @@ -0,0 +1,14 @@ +"""Deprecated: use splunk_ao.agent_stream instead of splunk_ao.__future__.agent_stream.""" + +import warnings + +warnings.warn( + "Importing from splunk_ao.__future__.agent_stream is deprecated. " + "Use splunk_ao.agent_stream instead.", + DeprecationWarning, + stacklevel=2, +) + +from splunk_ao.agent_stream import AgentStream # noqa: E402 + +__all__ = ["AgentStream"] diff --git a/src/splunk_ao/__future__/evaluator.py b/src/splunk_ao/__future__/evaluator.py new file mode 100644 index 00000000..bb7c5507 --- /dev/null +++ b/src/splunk_ao/__future__/evaluator.py @@ -0,0 +1,20 @@ +"""Deprecated: use splunk_ao.evaluator instead of splunk_ao.__future__.evaluator.""" + +import warnings + +warnings.warn( + "Importing from splunk_ao.__future__.evaluator is deprecated. " + "Use splunk_ao.evaluator instead.", + DeprecationWarning, + stacklevel=2, +) + +from splunk_ao.evaluator import ( # noqa: E402 + CodeEvaluator, + Evaluator, + LlmEvaluator, + LocalEvaluator, + SplunkAOEvaluator, +) + +__all__ = ["CodeEvaluator", "Evaluator", "LlmEvaluator", "LocalEvaluator", "SplunkAOEvaluator"] diff --git a/src/splunk_ao/__init__.py b/src/splunk_ao/__init__.py index 38b0611c..f0047cff 100644 --- a/src/splunk_ao/__init__.py +++ b/src/splunk_ao/__init__.py @@ -16,10 +16,18 @@ from galileo_core.schemas.logging.step import StepType from galileo_core.schemas.logging.trace import Trace from splunk_ao.agent_control import AgentControlTarget, AgentControlTargetUnresolvedError, get_agent_control_target +from splunk_ao.agent_stream import AgentStream from splunk_ao.collaborator import Collaborator, CollaboratorRole from splunk_ao.configuration import Configuration from splunk_ao.dataset import Dataset from splunk_ao.decorator import SplunkAODecorator, log, splunk_ao_context, start_session +from splunk_ao.evaluator import ( + CodeEvaluator, + Evaluator, + LlmEvaluator, + LocalEvaluator, + SplunkAOEvaluator, +) from splunk_ao.exceptions import ( AuthenticationError, BadRequestError, @@ -34,10 +42,8 @@ from splunk_ao.experiment import Experiment from splunk_ao.handlers.agent_control import SplunkAOAgentControlBridge, setup_agent_control_bridge from splunk_ao.integration import Integration -from splunk_ao.log_stream import LogStream from splunk_ao.logger import SplunkAOLogger from splunk_ao.logger.control import ControlAppliesTo, ControlCheckStage, ControlResult, ControlSpan -from splunk_ao.metric import CodeMetric, LlmMetric, LocalMetric, Metric, SplunkAOMetric from splunk_ao.model import Model from splunk_ao.project import Project from splunk_ao.prompt import Prompt @@ -64,6 +70,14 @@ __version__ = "0.1.0" __all__ = [ + # New canonical names (HYBIM-730) + "AgentStream", + "CodeEvaluator", + "Evaluator", + "LlmEvaluator", + "LocalEvaluator", + "SplunkAOEvaluator", + # Stable / unchanged "APIError", "AgentControlTarget", "AgentControlTargetUnresolvedError", @@ -74,7 +88,6 @@ "AzureProvider", "BadRequestError", "BedrockProvider", - "CodeMetric", "Collaborator", "CollaboratorRole", "Configuration", @@ -89,13 +102,9 @@ "Experiment", "ForbiddenError", "Integration", - "LlmMetric", "LlmSpan", - "LocalMetric", - "LogStream", "Message", "MessageRole", - "Metric", "MetricSpec", "MissingConfigurationError", "Model", @@ -118,7 +127,6 @@ "SplunkAOFutureError", "SplunkAOLogger", "SplunkAOLoggerException", - "SplunkAOMetric", "SplunkAOMetrics", "StepType", "StepWithChildSpans", @@ -141,3 +149,42 @@ "splunk_ao_context", "start_session", ] + +# --------------------------------------------------------------------------- +# PEP 562 — deprecated names emit DeprecationWarning on attribute access. +# Old names are NOT in __all__ so they won't appear in tab-completion, but +# they still work for existing code via __getattr__. +# --------------------------------------------------------------------------- +_DEPRECATED_NAMES: dict[str, tuple[str, str]] = { + "LogStream": ("AgentStream", "splunk_ao.agent_stream"), + "Metric": ("Evaluator", "splunk_ao.evaluator"), + "LlmMetric": ("LlmEvaluator", "splunk_ao.evaluator"), + "CodeMetric": ("CodeEvaluator", "splunk_ao.evaluator"), + "LocalMetric": ("LocalEvaluator", "splunk_ao.evaluator"), + "SplunkAOMetric": ("SplunkAOEvaluator", "splunk_ao.evaluator"), +} + +_DEPRECATED_OBJECTS: dict[str, object] = { + "LogStream": AgentStream, + "Metric": Evaluator, + "LlmMetric": LlmEvaluator, + "CodeMetric": CodeEvaluator, + "LocalMetric": LocalEvaluator, + "SplunkAOMetric": SplunkAOEvaluator, +} + + +def __getattr__(name: str) -> object: + if name in _DEPRECATED_NAMES: + import warnings + + new_name, new_module = _DEPRECATED_NAMES[name] + warnings.warn( + f"'splunk_ao.{name}' is deprecated and will be removed in a future release. " + f"Use '{new_name}' from '{new_module}' " + f"(or 'from splunk_ao import {new_name}') instead.", + DeprecationWarning, + stacklevel=2, + ) + return _DEPRECATED_OBJECTS[name] + raise AttributeError(f"module 'splunk_ao' has no attribute {name!r}") diff --git a/src/splunk_ao/agent_stream.py b/src/splunk_ao/agent_stream.py new file mode 100644 index 00000000..8a9f875b --- /dev/null +++ b/src/splunk_ao/agent_stream.py @@ -0,0 +1,61 @@ +""" +Agent Streams — the renamed successor to Log Streams (HYBIM-730). + +``AgentStream`` is the canonical name going forward. ``LogStream`` is kept +as a deprecated alias in ``splunk_ao.__init__`` and will be removed in a +future major release. + +The underlying implementation is shared with ``splunk_ao.log_stream``; the +API endpoints still use the ``/log_streams`` path (server-side rename is +tracked separately). +""" +from __future__ import annotations + +import warnings + +from splunk_ao.log_stream import LogStream + +__all__ = ["AgentStream"] + + +class AgentStream(LogStream): + """ + Object-centric interface for Splunk AO agent streams. + + ``AgentStream`` is the new name for what was previously called a + *Log Stream*. All functionality is identical; only the name has + changed. Use ``AgentStream`` for all new code. + + See ``splunk_ao.log_stream.LogStream`` for the full API reference — + every method, property, and class-method is inherited unchanged. + + Examples + -------- + from splunk_ao import AgentStream + + # Create and persist a new agent stream + stream = AgentStream(name="prod-traces", project_name="my-project").create() + + # Retrieve an existing stream + stream = AgentStream.get(name="prod-traces", project_name="my-project") + + # List streams for a project + streams = AgentStream.list(project_name="my-project") + + # Enable evaluators on the stream + from splunk_ao import SplunkAOMetrics + stream.set_metrics([SplunkAOMetrics.correctness, SplunkAOMetrics.completeness]) + """ + + +# Convenience: expose the deprecated ``LogStream`` name with a warning +def __getattr__(name: str): + if name == "LogStream": + warnings.warn( + "splunk_ao.agent_stream.LogStream is deprecated; " + "import AgentStream from splunk_ao.agent_stream instead.", + DeprecationWarning, + stacklevel=2, + ) + return LogStream + raise AttributeError(f"module 'splunk_ao.agent_stream' has no attribute {name!r}") diff --git a/src/splunk_ao/agent_streams.py b/src/splunk_ao/agent_streams.py new file mode 100644 index 00000000..a104e43b --- /dev/null +++ b/src/splunk_ao/agent_streams.py @@ -0,0 +1,208 @@ +""" +Agent Streams service layer — the renamed successor to Log Streams (HYBIM-730). + +Provides ``AgentStreams`` (service class) and the module-level convenience +functions ``get_agent_stream``, ``list_agent_streams``, ``create_agent_stream``, +and ``enable_evaluators``. + +The old ``log_streams`` module and its symbols remain available but are +deprecated. +""" +from __future__ import annotations + +import builtins +import warnings + +from splunk_ao.agent_stream import AgentStream +from splunk_ao.log_streams import LogStreams, create_log_stream, get_log_stream, list_log_streams +from splunk_ao.resources.types import Unset +from splunk_ao.schema.metrics import LocalMetricConfig, Metric, SplunkAOMetrics + +__all__ = [ + "AgentStreams", + "create_agent_stream", + "enable_evaluators", + "get_agent_stream", + "list_agent_streams", +] + + +class AgentStreams(LogStreams): + """ + Low-level service class for managing agent streams. + + ``AgentStreams`` is the new name for the ``LogStreams`` service class. + Inherits all methods from ``LogStreams`` and returns ``AgentStream`` + instances from listing/retrieval methods. + + Examples + -------- + from splunk_ao.agent_streams import AgentStreams + + svc = AgentStreams() + stream = svc.get(name="prod-traces", project_name="my-project") + streams = svc.list(project_name="my-project") + """ + + +# --------------------------------------------------------------------------- +# Module-level convenience functions +# --------------------------------------------------------------------------- + + +def get_agent_stream( + *, + name: str | None = None, + project_id: str | None = None, + project_name: str | None = None, +) -> AgentStream | None: + """ + Retrieve an agent stream by name. + + Parameters + ---------- + name: + The agent stream name. + project_id: + The project ID (mutually exclusive with *project_name*). + project_name: + The project name (mutually exclusive with *project_id*). + + Returns + ------- + AgentStream | None + The agent stream if found, ``None`` otherwise. + """ + result = get_log_stream(name=name, project_id=project_id, project_name=project_name) + if result is None: + return None + stream = AgentStream.__new__(AgentStream) + stream.__dict__.update(result.__dict__) + return stream + + +def list_agent_streams( + *, + project_id: str | None = None, + project_name: str | None = None, + limit: Unset | int = 100, + starting_token: Unset | int = 0, +) -> builtins.list[AgentStream]: + """ + List agent streams for a project. + + Parameters + ---------- + project_id: + The project ID (mutually exclusive with *project_name*). + project_name: + The project name (mutually exclusive with *project_id*). + limit: + Maximum number of results per page. Defaults to 100. + starting_token: + Pagination token. Defaults to 0. + + Returns + ------- + list[AgentStream] + A page of agent streams. + """ + log_streams = list_log_streams( + project_id=project_id, + project_name=project_name, + limit=limit, + starting_token=starting_token, + ) + result: builtins.list[AgentStream] = [] + for ls in log_streams: + s = AgentStream.__new__(AgentStream) + s.__dict__.update(ls.__dict__) + result.append(s) + return result + + +def create_agent_stream( + name: str, + project_id: str | None = None, + project_name: str | None = None, +) -> AgentStream: + """ + Create a new agent stream. + + Parameters + ---------- + name: + The agent stream name. + project_id: + The project ID (mutually exclusive with *project_name*). + project_name: + The project name (mutually exclusive with *project_id*). + + Returns + ------- + AgentStream + The created agent stream. + """ + ls = create_log_stream(name=name, project_id=project_id, project_name=project_name) + stream = AgentStream.__new__(AgentStream) + stream.__dict__.update(ls.__dict__) + return stream + + +def enable_evaluators( + *, + agent_stream_name: str | None = None, + project_name: str | None = None, + metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str], +) -> builtins.list[LocalMetricConfig]: + """ + Enable evaluators (formerly *metrics*) on an agent stream. + + Falls back to the ``SPLUNK_AO_LOG_STREAM`` and ``SPLUNK_AO_PROJECT`` + environment variables when *agent_stream_name* / *project_name* are + not provided explicitly. + + Parameters + ---------- + agent_stream_name: + The agent stream name. Falls back to ``SPLUNK_AO_LOG_STREAM`` env var. + project_name: + The project name. Falls back to ``SPLUNK_AO_PROJECT`` env var. + metrics: + Evaluators to enable. Accepts ``SplunkAOMetrics`` enum values, + ``Metric`` objects, ``LocalMetricConfig`` objects, or string names. + + Returns + ------- + list[LocalMetricConfig] + Local evaluator configurations that must be computed client-side. + """ + from splunk_ao.log_streams import LogStreams + + return LogStreams().enable_metrics( + log_stream_name=agent_stream_name, + project_name=project_name, + metrics=metrics, + ) + + +# --------------------------------------------------------------------------- +# Deprecated aliases for the old ``log_streams`` convenience functions +# --------------------------------------------------------------------------- + +def __getattr__(name: str): + _deprecated = { + "get_log_stream": ("get_agent_stream", get_agent_stream), + "list_log_streams": ("list_agent_streams", list_agent_streams), + "create_log_stream": ("create_agent_stream", create_agent_stream), + "enable_metrics": ("enable_evaluators", enable_evaluators), + } + if name in _deprecated: + new_name, obj = _deprecated[name] + warnings.warn( + f"splunk_ao.agent_streams.{name} is deprecated; use {new_name} instead.", + DeprecationWarning, + stacklevel=2, + ) + return obj + raise AttributeError(f"module 'splunk_ao.agent_streams' has no attribute {name!r}") diff --git a/src/splunk_ao/evaluator.py b/src/splunk_ao/evaluator.py new file mode 100644 index 00000000..b867c2a7 --- /dev/null +++ b/src/splunk_ao/evaluator.py @@ -0,0 +1,188 @@ +""" +Evaluators — the renamed successor to Metrics (HYBIM-730). + +``Evaluator`` and its concrete subclasses (``LlmEvaluator``, ``CodeEvaluator``, +``LocalEvaluator``, ``SplunkAOEvaluator``) are the canonical names going +forward. The old ``Metric``-prefixed names are kept as deprecated aliases in +``splunk_ao.__init__`` and will be removed in a future major release. + +The underlying implementation lives in ``splunk_ao.metric``; the API endpoints +still use the ``/scorers`` path (server-side rename is tracked separately). +""" +from __future__ import annotations + +import warnings + +from splunk_ao.metric import ( + BuiltInMetrics, + CodeMetric, + LlmMetric, + LocalMetric, + Metric, + SplunkAOMetric, +) + +__all__ = [ + "BuiltInEvaluators", + "CodeEvaluator", + "Evaluator", + "LlmEvaluator", + "LocalEvaluator", + "SplunkAOEvaluator", +] + + +class BuiltInEvaluators(BuiltInMetrics): + """ + Provides attribute-style access to built-in Splunk AO evaluators. + + This is the renamed successor to ``BuiltInMetrics``. Access built-in + evaluators via ``Evaluator.evaluators``. + + Examples + -------- + from splunk_ao import Evaluator + Evaluator.evaluators.correctness + Evaluator.evaluators.completeness + """ + + +class Evaluator(Metric): + """ + Base class for all Splunk AO evaluators. + + ``Evaluator`` is the new name for what was previously called a *Metric*. + All functionality is inherited from ``splunk_ao.metric.Metric`` unchanged. + + Use one of the concrete subclasses for new code: + + - ``SplunkAOEvaluator`` — built-in Splunk AO scorers (``Evaluator.evaluators.*``) + - ``LlmEvaluator`` — custom LLM-judge evaluators + - ``LocalEvaluator`` — local function-based evaluators + - ``CodeEvaluator`` — code-based evaluators + + Class Attributes + ---------------- + evaluators : BuiltInEvaluators + Access built-in Splunk AO evaluators. + + Examples + -------- + from splunk_ao import Evaluator, AgentStream, SplunkAOMetrics + + stream = AgentStream.get(name="prod-traces", project_name="my-project") + stream.set_metrics([ + Evaluator.evaluators.correctness, + Evaluator.evaluators.completeness, + ]) + + # Get an existing evaluator by name + ev = Evaluator.get(name="factuality-checker") + + # List all evaluators + evaluators = Evaluator.list() + """ + + evaluators = BuiltInEvaluators() + + # Keep ``metrics`` attribute as deprecated alias + @property # type: ignore[override] + def metrics(self): # type: ignore[override] + warnings.warn( + "Evaluator.metrics is deprecated; use Evaluator.evaluators instead.", + DeprecationWarning, + stacklevel=2, + ) + return self.__class__.evaluators + + +class LlmEvaluator(LlmMetric): + """ + LLM-based evaluator — the renamed successor to ``LlmMetric``. + + See ``splunk_ao.metric.LlmMetric`` for the full API reference. + + Examples + -------- + from splunk_ao import LlmEvaluator + + ev = LlmEvaluator( + name="response_quality", + prompt="Rate the quality 1-10: {input} -> {output}", + model="gpt-4o-mini", + judges=3, + ).create() + """ + + +class CodeEvaluator(CodeMetric): + """ + Code-based evaluator — the renamed successor to ``CodeMetric``. + + See ``splunk_ao.metric.CodeMetric`` for the full API reference. + + Examples + -------- + from splunk_ao import CodeEvaluator + + ev = CodeEvaluator( + name="custom_scorer", + code="def scorer_fn(step): return 1.0", + ).create() + """ + + +class LocalEvaluator(LocalMetric): + """ + Local function-based evaluator — the renamed successor to ``LocalMetric``. + + See ``splunk_ao.metric.LocalMetric`` for the full API reference. + + Examples + -------- + from splunk_ao import LocalEvaluator + + def my_fn(trace): + return 0.5 + + ev = LocalEvaluator(name="my_scorer", scorer_fn=my_fn) + """ + + +class SplunkAOEvaluator(SplunkAOMetric): + """ + Built-in Splunk AO scorer evaluator — the renamed successor to ``SplunkAOMetric``. + + Access built-in evaluators via ``Evaluator.evaluators.*``. + + Examples + -------- + from splunk_ao import Evaluator + + ev = Evaluator.get(name="correctness") + assert isinstance(ev, SplunkAOEvaluator) + """ + + +# --------------------------------------------------------------------------- +# Deprecated aliases for the old Metric-prefixed names +# --------------------------------------------------------------------------- + +def __getattr__(name: str): + _deprecated = { + "Metric": ("Evaluator", Evaluator), + "LlmMetric": ("LlmEvaluator", LlmEvaluator), + "CodeMetric": ("CodeEvaluator", CodeEvaluator), + "LocalMetric": ("LocalEvaluator", LocalEvaluator), + "SplunkAOMetric": ("SplunkAOEvaluator", SplunkAOEvaluator), + "BuiltInMetrics": ("BuiltInEvaluators", BuiltInEvaluators), + } + if name in _deprecated: + new_name, obj = _deprecated[name] + warnings.warn( + f"splunk_ao.evaluator.{name} is deprecated; use {new_name} instead.", + DeprecationWarning, + stacklevel=2, + ) + return obj + raise AttributeError(f"module 'splunk_ao.evaluator' has no attribute {name!r}") diff --git a/src/splunk_ao/evaluators.py b/src/splunk_ao/evaluators.py new file mode 100644 index 00000000..25b6554d --- /dev/null +++ b/src/splunk_ao/evaluators.py @@ -0,0 +1,194 @@ +""" +Evaluators service layer — the renamed successor to Metrics (HYBIM-730). + +Provides ``Evaluators`` (service class) and the module-level convenience +functions ``create_custom_llm_evaluator``, ``get_evaluators``, and +``delete_evaluator``. + +The old ``metrics`` module and its symbols remain available but are deprecated. +""" +from __future__ import annotations + +import datetime +import warnings + +from splunk_ao.metrics import Metrics, create_custom_llm_metric, delete_metric, get_metrics +from splunk_ao.resources.models.base_scorer_version_response import BaseScorerVersionResponse +from splunk_ao.resources.models.output_type_enum import OutputTypeEnum +from splunk_ao.resources.models.log_records_metrics_response import LogRecordsMetricsResponse +from galileo_core.schemas.logging.step import StepType +from splunk_ao.search import FilterType + +__all__ = [ + "Evaluators", + "create_custom_llm_evaluator", + "delete_evaluator", + "get_evaluators", +] + + +class Evaluators(Metrics): + """ + Low-level service class for managing evaluators. + + ``Evaluators`` is the new name for the ``Metrics`` service class. + Inherits all methods from ``Metrics`` unchanged. + + Examples + -------- + from splunk_ao.evaluators import Evaluators + + svc = Evaluators() + svc.delete_metric(name="old-evaluator") + """ + + +# --------------------------------------------------------------------------- +# Module-level convenience functions +# --------------------------------------------------------------------------- + + +def create_custom_llm_evaluator( + name: str, + user_prompt: str, + node_level: StepType = StepType.llm, + cot_enabled: bool = True, + model_name: str = "gpt-4.1-mini", + num_judges: int = 3, + description: str = "", + tags: list[str] | None = None, + output_type: OutputTypeEnum = OutputTypeEnum.BOOLEAN, + ground_truth: bool = False, +) -> BaseScorerVersionResponse: + """ + Create a custom LLM evaluator. + + This is the renamed equivalent of ``create_custom_llm_metric`` from + ``splunk_ao.metrics``. + + Parameters + ---------- + name: + Name of the evaluator. + user_prompt: + Prompt template for the evaluator. + node_level: + Node level. Defaults to ``StepType.llm``. + cot_enabled: + Whether chain-of-thought reasoning is enabled. + model_name: + Model alias to use for judging. + num_judges: + Number of judge LLMs to use. + description: + Human-readable description. + tags: + Tags to associate with the evaluator. + output_type: + Output type (boolean, percentage, etc.). + ground_truth: + Whether the evaluator requires a ground-truth reference value. + + Returns + ------- + BaseScorerVersionResponse + The created evaluator version details. + """ + return create_custom_llm_metric( + name=name, + user_prompt=user_prompt, + node_level=node_level, + cot_enabled=cot_enabled, + model_name=model_name, + num_judges=num_judges, + description=description, + tags=tags, + output_type=output_type, + ground_truth=ground_truth, + ) + + +def get_evaluators( + project_id: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + experiment_id: str | None = None, + agent_stream_id: str | None = None, + filters: list[FilterType] | None = None, + group_by: str | None = None, + interval: int = 5, +) -> LogRecordsMetricsResponse: + """ + Query evaluator results for a project. + + This is the renamed equivalent of ``get_metrics`` from ``splunk_ao.metrics``. + + Parameters + ---------- + project_id: + Project UUID. + start_time: + Start of the query window. + end_time: + End of the query window. + experiment_id: + Filter by experiment ID (optional). + agent_stream_id: + Filter by agent stream ID (optional). + filters: + Additional query filters. + group_by: + Field to group results by. + interval: + Time interval in seconds. + + Returns + ------- + LogRecordsMetricsResponse + Evaluator query results. + """ + return get_metrics( + project_id=project_id, + start_time=start_time, + end_time=end_time, + experiment_id=experiment_id, + log_stream_id=agent_stream_id, + filters=filters, + group_by=group_by, + interval=interval, + ) + + +def delete_evaluator(name: str) -> None: + """ + Delete an evaluator by name. + + This is the renamed equivalent of ``delete_metric`` from ``splunk_ao.metrics``. + + Parameters + ---------- + name: + The evaluator name to delete. + """ + return delete_metric(name=name) + + +# --------------------------------------------------------------------------- +# Deprecated aliases for old ``metrics`` module function names +# --------------------------------------------------------------------------- + +def __getattr__(name: str): + _deprecated = { + "create_custom_llm_metric": ("create_custom_llm_evaluator", create_custom_llm_evaluator), + "get_metrics": ("get_evaluators", get_evaluators), + "delete_metric": ("delete_evaluator", delete_evaluator), + } + if name in _deprecated: + new_name, obj = _deprecated[name] + warnings.warn( + f"splunk_ao.evaluators.{name} is deprecated; use {new_name} instead.", + DeprecationWarning, + stacklevel=2, + ) + return obj + raise AttributeError(f"module 'splunk_ao.evaluators' has no attribute {name!r}") diff --git a/src/splunk_ao/project.py b/src/splunk_ao/project.py index 05ca34dd..4b763d92 100644 --- a/src/splunk_ao/project.py +++ b/src/splunk_ao/project.py @@ -288,64 +288,114 @@ def list(cls) -> builtins.list[Project]: return [cls._from_api_response(retrieved_project) for retrieved_project in retrieved_projects] - def create_log_stream(self, name: str) -> LogStream: + def create_agent_stream(self, name: str) -> "AgentStream": """ - Create a new log stream for this project. + Create a new agent stream for this project. Args: - name (str): The name of the log stream to create. + name (str): The name of the agent stream to create. Returns ------- - LogStream: The created log stream. + AgentStream: The created agent stream. Examples -------- project = Project.get(name="My AI Project") - log_stream = project.create_log_stream(name="Production Logs") + stream = project.create_agent_stream(name="Production Traces") """ + from splunk_ao.agent_stream import AgentStream # lazy to avoid circular import + if self.id is None: - raise ValueError("Project ID is not set. Cannot create log stream for a local-only project.") + raise ValueError("Project ID is not set. Cannot create agent stream for a local-only project.") - # Use the LogStream pattern to avoid duplication - return LogStream(name=name, project_id=self.id).create() + return AgentStream(name=name, project_id=self.id).create() - def list_log_streams( + def create_log_stream(self, name: str) -> "AgentStream": + """ + Create a new agent stream for this project. + + .. deprecated:: + Use :meth:`create_agent_stream` instead. + + Args: + name (str): The name of the agent stream to create. + + Returns + ------- + AgentStream: The created agent stream. + + Examples + -------- + project = Project.get(name="My AI Project") + log_stream = project.create_log_stream(name="Production Logs") + """ + import warnings + warnings.warn( + "Project.create_log_stream() is deprecated; use Project.create_agent_stream() instead.", + DeprecationWarning, + stacklevel=2, + ) + return self.create_agent_stream(name=name) + + def list_agent_streams( self, *, limit: Unset | int = 100, starting_token: Unset | int = 0 - ) -> builtins.list[LogStream]: + ) -> "builtins.list[AgentStream]": """ - List log streams for this project. + List agent streams for this project. Returns a single page of results. Use `starting_token` (from `next_starting_token` on a prior response) to fetch subsequent pages. Args: - limit (Union[Unset, int]): Maximum number of log streams to return per page. Defaults to 100. - starting_token (Union[Unset, int]): Pagination token to start from. Defaults to 0 (first page). + limit (Union[Unset, int]): Maximum number of agent streams per page. Defaults to 100. + starting_token (Union[Unset, int]): Pagination token. Defaults to 0 (first page). Returns ------- - List[LogStream]: A page of log streams belonging to this project. + List[AgentStream]: A page of agent streams belonging to this project. Examples -------- project = Project.get(name="My AI Project") - log_streams = project.list_log_streams() - for stream in log_streams: - # Process each log stream + streams = project.list_agent_streams() + for stream in streams: pass + """ + from splunk_ao.agent_stream import AgentStream # lazy to avoid circular import - # Cap the number of returned log streams - log_streams = project.list_log_streams(limit=3) + if self.id is None: + raise ValueError("Project ID is not set. Cannot list agent streams for a local-only project.") - # Fetch the next page - page_2 = project.list_log_streams(starting_token=100) + return AgentStream.list(project_id=self.id, limit=limit, starting_token=starting_token) + + def list_log_streams( + self, *, limit: Unset | int = 100, starting_token: Unset | int = 0 + ) -> "builtins.list[AgentStream]": """ - if self.id is None: - raise ValueError("Project ID is not set. Cannot list log streams for a local-only project.") + List agent streams for this project. + + .. deprecated:: + Use :meth:`list_agent_streams` instead. + + Returns a single page of results. Use `starting_token` (from + `next_starting_token` on a prior response) to fetch subsequent pages. - # Use the LogStream pattern to avoid duplication - return LogStream.list(project_id=self.id, limit=limit, starting_token=starting_token) + Args: + limit (Union[Unset, int]): Maximum number of streams per page. Defaults to 100. + starting_token (Union[Unset, int]): Pagination token. Defaults to 0 (first page). + + Returns + ------- + List[AgentStream]: A page of agent streams belonging to this project. + """ + import warnings + warnings.warn( + "Project.list_log_streams() is deprecated; use Project.list_agent_streams() instead.", + DeprecationWarning, + stacklevel=2, + ) + return self.list_agent_streams(limit=limit, starting_token=starting_token) def list_experiments(self) -> builtins.list[Experiment]: """ @@ -414,24 +464,44 @@ def list_prompts(self) -> builtins.list[Prompt]: return Prompt.list(project_id=self.id) @property - def logstreams(self) -> builtins.list[LogStream]: + def agent_streams(self) -> "builtins.list[AgentStream]": """ - Property to access log streams for this project. + Property to access agent streams for this project. - This is a read-only property that returns the current list of log streams. - To create new log streams, use create_log_stream(). + This is a read-only property that returns the current list of agent streams. + To create new agent streams, use :meth:`create_agent_stream`. Returns ------- - List[LogStream]: A list of log streams belonging to this project. + List[AgentStream]: A list of agent streams belonging to this project. Examples -------- project = Project.get(name="My AI Project") - for stream in project.logstreams: + for stream in project.agent_streams: print(stream.name) """ - return self.list_log_streams() + return self.list_agent_streams() + + @property + def logstreams(self) -> "builtins.list[AgentStream]": + """ + Property to access agent streams for this project. + + .. deprecated:: + Use :attr:`agent_streams` instead. + + Returns + ------- + List[AgentStream]: A list of agent streams belonging to this project. + """ + import warnings + warnings.warn( + "Project.logstreams is deprecated; use Project.agent_streams instead.", + DeprecationWarning, + stacklevel=2, + ) + return self.list_agent_streams() @property def experiments(self) -> builtins.list[Experiment]: @@ -895,3 +965,6 @@ def save(self) -> Project: from splunk_ao.experiment import Experiment # noqa: E402 from splunk_ao.log_stream import LogStream # noqa: E402 from splunk_ao.prompt import Prompt # noqa: E402 + +# AgentStream is imported lazily inside methods to avoid a secondary circular import: +# agent_stream → log_stream → project → agent_stream From 74679dc7a74eec60bffe160aac94f01d8f3bc787 Mon Sep 17 00:00:00 2001 From: adityamehra Date: Mon, 20 Jul 2026 09:28:35 -0700 Subject: [PATCH 2/4] fix: resolve mypy type errors in domain rename modules - Add return type -> object to all __getattr__ module functions - Add AgentStream to TYPE_CHECKING import in project.py so string annotations resolve correctly under mypy - Remove incorrect # type: ignore[override] from Evaluator.metrics property and add explicit return type BuiltInEvaluators Co-authored-by: Cursor --- src/splunk_ao/agent_stream.py | 2 +- src/splunk_ao/agent_streams.py | 2 +- src/splunk_ao/evaluator.py | 6 +++--- src/splunk_ao/evaluators.py | 2 +- src/splunk_ao/project.py | 1 + 5 files changed, 7 insertions(+), 6 deletions(-) diff --git a/src/splunk_ao/agent_stream.py b/src/splunk_ao/agent_stream.py index 8a9f875b..127a0156 100644 --- a/src/splunk_ao/agent_stream.py +++ b/src/splunk_ao/agent_stream.py @@ -49,7 +49,7 @@ class AgentStream(LogStream): # Convenience: expose the deprecated ``LogStream`` name with a warning -def __getattr__(name: str): +def __getattr__(name: str) -> object: if name == "LogStream": warnings.warn( "splunk_ao.agent_stream.LogStream is deprecated; " diff --git a/src/splunk_ao/agent_streams.py b/src/splunk_ao/agent_streams.py index a104e43b..28484c7e 100644 --- a/src/splunk_ao/agent_streams.py +++ b/src/splunk_ao/agent_streams.py @@ -190,7 +190,7 @@ def enable_evaluators( # Deprecated aliases for the old ``log_streams`` convenience functions # --------------------------------------------------------------------------- -def __getattr__(name: str): +def __getattr__(name: str) -> object: _deprecated = { "get_log_stream": ("get_agent_stream", get_agent_stream), "list_log_streams": ("list_agent_streams", list_agent_streams), diff --git a/src/splunk_ao/evaluator.py b/src/splunk_ao/evaluator.py index b867c2a7..f1916f54 100644 --- a/src/splunk_ao/evaluator.py +++ b/src/splunk_ao/evaluator.py @@ -86,8 +86,8 @@ class Evaluator(Metric): evaluators = BuiltInEvaluators() # Keep ``metrics`` attribute as deprecated alias - @property # type: ignore[override] - def metrics(self): # type: ignore[override] + @property + def metrics(self) -> BuiltInEvaluators: warnings.warn( "Evaluator.metrics is deprecated; use Evaluator.evaluators instead.", DeprecationWarning, @@ -168,7 +168,7 @@ class SplunkAOEvaluator(SplunkAOMetric): # Deprecated aliases for the old Metric-prefixed names # --------------------------------------------------------------------------- -def __getattr__(name: str): +def __getattr__(name: str) -> object: _deprecated = { "Metric": ("Evaluator", Evaluator), "LlmMetric": ("LlmEvaluator", LlmEvaluator), diff --git a/src/splunk_ao/evaluators.py b/src/splunk_ao/evaluators.py index 25b6554d..fcec5bed 100644 --- a/src/splunk_ao/evaluators.py +++ b/src/splunk_ao/evaluators.py @@ -177,7 +177,7 @@ def delete_evaluator(name: str) -> None: # Deprecated aliases for old ``metrics`` module function names # --------------------------------------------------------------------------- -def __getattr__(name: str): +def __getattr__(name: str) -> object: _deprecated = { "create_custom_llm_metric": ("create_custom_llm_evaluator", create_custom_llm_evaluator), "get_metrics": ("get_evaluators", get_evaluators), diff --git a/src/splunk_ao/project.py b/src/splunk_ao/project.py index 4b763d92..497623bf 100644 --- a/src/splunk_ao/project.py +++ b/src/splunk_ao/project.py @@ -16,6 +16,7 @@ from splunk_ao.shared.exceptions import APIError, ValidationError if TYPE_CHECKING: + from splunk_ao.agent_stream import AgentStream from splunk_ao.dataset import Dataset from splunk_ao.experiment import Experiment from splunk_ao.log_stream import LogStream From 3aee54dbc1db4622ce659fc4f58ccb25667eece4 Mon Sep 17 00:00:00 2001 From: adityamehra Date: Mon, 20 Jul 2026 10:29:17 -0700 Subject: [PATCH 3/4] fix: resolve remaining mypy override errors in domain rename - evaluator.py: remove deprecated `metrics` property; mypy forbids overriding a writable class attribute (Metric.metrics = BuiltInMetrics()) with a read-only property in a subclass. Evaluator.evaluators remains the canonical accessor. - project.py: add cast() around AgentStream.create() and AgentStream.list() return values. Both methods inherit LogStream's signature (-> LogStream / -> list[LogStream]); cast tells mypy the runtime types are AgentStream. Co-authored-by: Cursor --- src/splunk_ao/evaluator.py | 11 ++--------- src/splunk_ao/project.py | 10 +++++++--- 2 files changed, 9 insertions(+), 12 deletions(-) diff --git a/src/splunk_ao/evaluator.py b/src/splunk_ao/evaluator.py index f1916f54..b9b3e360 100644 --- a/src/splunk_ao/evaluator.py +++ b/src/splunk_ao/evaluator.py @@ -85,15 +85,8 @@ class Evaluator(Metric): evaluators = BuiltInEvaluators() - # Keep ``metrics`` attribute as deprecated alias - @property - def metrics(self) -> BuiltInEvaluators: - warnings.warn( - "Evaluator.metrics is deprecated; use Evaluator.evaluators instead.", - DeprecationWarning, - stacklevel=2, - ) - return self.__class__.evaluators + # ``metrics`` class attribute is inherited from Metric and intentionally + # left as-is. The new canonical accessor is ``Evaluator.evaluators``. class LlmEvaluator(LlmMetric): diff --git a/src/splunk_ao/project.py b/src/splunk_ao/project.py index 497623bf..9ae5e9e9 100644 --- a/src/splunk_ao/project.py +++ b/src/splunk_ao/project.py @@ -3,7 +3,7 @@ import builtins import logging from datetime import datetime -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, cast from splunk_ao.collaborator import Collaborator, CollaboratorRole from splunk_ao.config import SplunkAOConfig @@ -310,7 +310,9 @@ def create_agent_stream(self, name: str) -> "AgentStream": if self.id is None: raise ValueError("Project ID is not set. Cannot create agent stream for a local-only project.") - return AgentStream(name=name, project_id=self.id).create() + # cast: LogStream.create() returns LogStream, but self is AgentStream so + # the runtime type is correct. cast tells mypy to trust us. + return cast("AgentStream", AgentStream(name=name, project_id=self.id).create()) def create_log_stream(self, name: str) -> "AgentStream": """ @@ -368,7 +370,9 @@ def list_agent_streams( if self.id is None: raise ValueError("Project ID is not set. Cannot list agent streams for a local-only project.") - return AgentStream.list(project_id=self.id, limit=limit, starting_token=starting_token) + # cast: LogStream.list() is typed -> list[LogStream]; runtime values are + # AgentStream instances because cls is AgentStream. + return cast("builtins.list[AgentStream]", AgentStream.list(project_id=self.id, limit=limit, starting_token=starting_token)) def list_log_streams( self, *, limit: Unset | int = 100, starting_token: Unset | int = 0 From 684945e805aa8cfd559768c2972be449039f57e9 Mon Sep 17 00:00:00 2001 From: adityamehra Date: Tue, 21 Jul 2026 16:25:37 -0700 Subject: [PATCH 4/4] =?UTF-8?q?feat(HYBIM-730):=20hard=20cut-over=20?= =?UTF-8?q?=E2=80=94=20pure=20rename,=20no=20backward=20compat=20shims?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Domain entity renames (per PR review: no deprecated wrappers): LogStream → AgentStream - src/splunk_ao/log_stream.py → agent_stream.py (AgentStream class) - src/splunk_ao/log_streams.py → agent_streams.py (AgentStreams service) - enable_metrics() → enable_evaluators() - get_log_stream/list_log_streams/create_log_stream → get/list/create_agent_stream - project.py: remove deprecated create_log_stream/list_log_streams/logstreams Metric → Evaluator - src/splunk_ao/metric.py → evaluator.py (Evaluator, LlmEvaluator, CodeEvaluator, LocalEvaluator, SplunkAOEvaluator, BuiltInEvaluators) - src/splunk_ao/metrics.py → evaluators.py (Evaluators service, renamed convenience functions) - delete_metric → delete_evaluator - create_custom_llm_metric → create_custom_llm_evaluator - get_metrics → get_evaluators Other - __init__.py: remove PEP 562 deprecated __getattr__ shims; export only new names - __future__/__init__.py: updated to new class names - types.py: MetricSpec uses Evaluator - export.py, logger/logger.py: updated to AgentStreams - All test files updated (imports, @patch paths, assertions, param names) - Tests renamed: test_log_stream* → test_agent_stream*, test_metric* → test_evaluator* Results: 1856 passed, 18 failed (all pre-existing), 3 error (pre-existing) Type-check: Success (107 source files) Co-authored-by: Cursor --- src/splunk_ao/__future__/__init__.py | 17 +- src/splunk_ao/__future__/agent_stream.py | 14 - src/splunk_ao/__future__/evaluator.py | 20 - src/splunk_ao/__future__/log_stream.py | 13 - src/splunk_ao/__future__/metric.py | 13 - src/splunk_ao/__init__.py | 38 - src/splunk_ao/agent_stream.py | 990 +++++++++++- src/splunk_ao/agent_streams.py | 848 +++++++++-- src/splunk_ao/evaluator.py | 1350 +++++++++++++++-- src/splunk_ao/evaluators.py | 309 ++-- src/splunk_ao/export.py | 4 +- src/splunk_ao/log_stream.py | 971 ------------ src/splunk_ao/log_streams.py | 790 ---------- src/splunk_ao/logger/logger.py | 4 +- src/splunk_ao/metric.py | 1317 ---------------- src/splunk_ao/metrics.py | 259 ---- src/splunk_ao/project.py | 102 +- src/splunk_ao/types.py | 4 +- tests/test_agent_control_bridge.py | 24 +- ...est_log_stream.py => test_agent_stream.py} | 192 +-- ...cs.py => test_agent_streams_evaluators.py} | 110 +- ...on.py => test_agent_streams_pagination.py} | 104 +- tests/test_async_base_handler.py | 2 +- tests/test_backward_compat_future.py | 56 +- tests/test_base_handler.py | 2 +- tests/test_crewai_handler.py | 2 +- tests/test_decorator.py | 114 +- tests/test_decorator_distributed.py | 24 +- tests/{test_metric.py => test_evaluator.py} | 556 +++---- ...etric_types.py => test_evaluator_types.py} | 192 +-- tests/{test_metrics.py => test_evaluators.py} | 130 +- tests/test_experiments.py | 12 +- tests/test_export.py | 12 +- tests/test_galileo_context.py | 4 +- tests/test_langchain.py | 4 +- tests/test_langchain_async.py | 2 +- tests/test_langchain_middleware.py | 4 +- tests/test_logger_batch.py | 88 +- tests/test_logger_distributed.py | 74 +- tests/test_logger_timestamps.py | 6 +- tests/test_middleware_tracing.py | 14 +- tests/test_openai.py | 26 +- tests/test_openai_agents.py | 6 +- tests/testutils/setup.py | 6 +- 44 files changed, 4004 insertions(+), 4825 deletions(-) delete mode 100644 src/splunk_ao/__future__/agent_stream.py delete mode 100644 src/splunk_ao/__future__/evaluator.py delete mode 100644 src/splunk_ao/__future__/log_stream.py delete mode 100644 src/splunk_ao/__future__/metric.py delete mode 100644 src/splunk_ao/log_stream.py delete mode 100644 src/splunk_ao/log_streams.py delete mode 100644 src/splunk_ao/metric.py delete mode 100644 src/splunk_ao/metrics.py rename tests/{test_log_stream.py => test_agent_stream.py} (89%) rename tests/{test_log_streams_metrics.py => test_agent_streams_evaluators.py} (82%) rename tests/{test_log_streams_pagination.py => test_agent_streams_pagination.py} (71%) rename tests/{test_metric.py => test_evaluator.py} (77%) rename tests/{test_metric_types.py => test_evaluator_types.py} (60%) rename tests/{test_metrics.py => test_evaluators.py} (79%) diff --git a/src/splunk_ao/__future__/__init__.py b/src/splunk_ao/__future__/__init__.py index 0d50ff2a..efe7bae5 100644 --- a/src/splunk_ao/__future__/__init__.py +++ b/src/splunk_ao/__future__/__init__.py @@ -7,8 +7,8 @@ from splunk_ao.dataset import Dataset from splunk_ao.experiment import Experiment from splunk_ao.integration import Integration -from splunk_ao.log_stream import LogStream -from splunk_ao.metric import CodeMetric, LlmMetric, LocalMetric, Metric, SplunkAOMetric +from splunk_ao.agent_stream import AgentStream +from splunk_ao.evaluator import BuiltInEvaluators, CodeEvaluator, Evaluator, LlmEvaluator, LocalEvaluator, SplunkAOEvaluator from splunk_ao.model import Model from splunk_ao.project import Project from splunk_ao.prompt import Prompt @@ -26,28 +26,29 @@ __all__ = [ "APIError", - "CodeMetric", + "AgentStream", + "BuiltInEvaluators", + "CodeEvaluator", "Collaborator", "CollaboratorRole", "Configuration", "ConfigurationError", "Dataset", + "Evaluator", "Experiment", "Integration", - "LlmMetric", - "LocalMetric", - "LogStream", + "LlmEvaluator", + "LocalEvaluator", "Message", "MessageRole", - "Metric", "Model", "Project", "Prompt", "RecordType", "ResourceConflictError", "ResourceNotFoundError", + "SplunkAOEvaluator", "SplunkAOFutureError", - "SplunkAOMetric", "StepType", "ValidationError", "enable_console_logging", diff --git a/src/splunk_ao/__future__/agent_stream.py b/src/splunk_ao/__future__/agent_stream.py deleted file mode 100644 index 27a7e092..00000000 --- a/src/splunk_ao/__future__/agent_stream.py +++ /dev/null @@ -1,14 +0,0 @@ -"""Deprecated: use splunk_ao.agent_stream instead of splunk_ao.__future__.agent_stream.""" - -import warnings - -warnings.warn( - "Importing from splunk_ao.__future__.agent_stream is deprecated. " - "Use splunk_ao.agent_stream instead.", - DeprecationWarning, - stacklevel=2, -) - -from splunk_ao.agent_stream import AgentStream # noqa: E402 - -__all__ = ["AgentStream"] diff --git a/src/splunk_ao/__future__/evaluator.py b/src/splunk_ao/__future__/evaluator.py deleted file mode 100644 index bb7c5507..00000000 --- a/src/splunk_ao/__future__/evaluator.py +++ /dev/null @@ -1,20 +0,0 @@ -"""Deprecated: use splunk_ao.evaluator instead of splunk_ao.__future__.evaluator.""" - -import warnings - -warnings.warn( - "Importing from splunk_ao.__future__.evaluator is deprecated. " - "Use splunk_ao.evaluator instead.", - DeprecationWarning, - stacklevel=2, -) - -from splunk_ao.evaluator import ( # noqa: E402 - CodeEvaluator, - Evaluator, - LlmEvaluator, - LocalEvaluator, - SplunkAOEvaluator, -) - -__all__ = ["CodeEvaluator", "Evaluator", "LlmEvaluator", "LocalEvaluator", "SplunkAOEvaluator"] diff --git a/src/splunk_ao/__future__/log_stream.py b/src/splunk_ao/__future__/log_stream.py deleted file mode 100644 index 5c090f22..00000000 --- a/src/splunk_ao/__future__/log_stream.py +++ /dev/null @@ -1,13 +0,0 @@ -"""Deprecated: use splunk_ao.log_stream instead of splunk_ao.__future__.log_stream.""" - -import warnings - -warnings.warn( - "Importing from splunk_ao.__future__.log_stream is deprecated. Use splunk_ao.log_stream instead.", - DeprecationWarning, - stacklevel=2, -) - -from splunk_ao.log_stream import LogStream # noqa: E402 - -__all__ = ["LogStream"] diff --git a/src/splunk_ao/__future__/metric.py b/src/splunk_ao/__future__/metric.py deleted file mode 100644 index 3df9a658..00000000 --- a/src/splunk_ao/__future__/metric.py +++ /dev/null @@ -1,13 +0,0 @@ -"""Deprecated: use splunk_ao.metric instead of splunk_ao.__future__.metric.""" - -import warnings - -warnings.warn( - "Importing from splunk_ao.__future__.metric is deprecated. Use splunk_ao.metric instead.", - DeprecationWarning, - stacklevel=2, -) - -from splunk_ao.metric import BuiltInMetrics, CodeMetric, LlmMetric, LocalMetric, Metric, SplunkAOMetric # noqa: E402 - -__all__ = ["BuiltInMetrics", "CodeMetric", "LlmMetric", "LocalMetric", "Metric", "SplunkAOMetric"] diff --git a/src/splunk_ao/__init__.py b/src/splunk_ao/__init__.py index f0047cff..2c3a92e9 100644 --- a/src/splunk_ao/__init__.py +++ b/src/splunk_ao/__init__.py @@ -150,41 +150,3 @@ "start_session", ] -# --------------------------------------------------------------------------- -# PEP 562 — deprecated names emit DeprecationWarning on attribute access. -# Old names are NOT in __all__ so they won't appear in tab-completion, but -# they still work for existing code via __getattr__. -# --------------------------------------------------------------------------- -_DEPRECATED_NAMES: dict[str, tuple[str, str]] = { - "LogStream": ("AgentStream", "splunk_ao.agent_stream"), - "Metric": ("Evaluator", "splunk_ao.evaluator"), - "LlmMetric": ("LlmEvaluator", "splunk_ao.evaluator"), - "CodeMetric": ("CodeEvaluator", "splunk_ao.evaluator"), - "LocalMetric": ("LocalEvaluator", "splunk_ao.evaluator"), - "SplunkAOMetric": ("SplunkAOEvaluator", "splunk_ao.evaluator"), -} - -_DEPRECATED_OBJECTS: dict[str, object] = { - "LogStream": AgentStream, - "Metric": Evaluator, - "LlmMetric": LlmEvaluator, - "CodeMetric": CodeEvaluator, - "LocalMetric": LocalEvaluator, - "SplunkAOMetric": SplunkAOEvaluator, -} - - -def __getattr__(name: str) -> object: - if name in _DEPRECATED_NAMES: - import warnings - - new_name, new_module = _DEPRECATED_NAMES[name] - warnings.warn( - f"'splunk_ao.{name}' is deprecated and will be removed in a future release. " - f"Use '{new_name}' from '{new_module}' " - f"(or 'from splunk_ao import {new_name}') instead.", - DeprecationWarning, - stacklevel=2, - ) - return _DEPRECATED_OBJECTS[name] - raise AttributeError(f"module 'splunk_ao' has no attribute {name!r}") diff --git a/src/splunk_ao/agent_stream.py b/src/splunk_ao/agent_stream.py index 127a0156..a42af675 100644 --- a/src/splunk_ao/agent_stream.py +++ b/src/splunk_ao/agent_stream.py @@ -1,61 +1,973 @@ -""" -Agent Streams — the renamed successor to Log Streams (HYBIM-730). +from __future__ import annotations -``AgentStream`` is the canonical name going forward. ``LogStream`` is kept -as a deprecated alias in ``splunk_ao.__init__`` and will be removed in a -future major release. +import builtins +import logging +from collections.abc import Iterator +from datetime import datetime +from typing import TYPE_CHECKING, Any -The underlying implementation is shared with ``splunk_ao.log_stream``; the -API endpoints still use the ``/log_streams`` path (server-side rename is -tracked separately). -""" -from __future__ import annotations +from splunk_ao.config import SplunkAOConfig +from splunk_ao.decorator import splunk_ao_context +from splunk_ao.export import ExportClient +from splunk_ao.agent_streams import AgentStreams +from splunk_ao.resources.api.trace import ( + sessions_available_columns_projects_project_id_sessions_available_columns_post, + spans_available_columns_projects_project_id_spans_available_columns_post, + traces_available_columns_projects_project_id_traces_available_columns_post, +) +from splunk_ao.resources.models import LLMExportFormat, LogRecordsSortClause, RootType +from splunk_ao.resources.models.http_validation_error import HTTPValidationError +from splunk_ao.resources.models.log_records_available_columns_request import LogRecordsAvailableColumnsRequest +from splunk_ao.resources.models.log_records_available_columns_response import LogRecordsAvailableColumnsResponse +from splunk_ao.resources.types import Unset +from splunk_ao.schema.filters import FilterType +from splunk_ao.schema.metrics import LocalMetricConfig, Metric, SplunkAOMetrics +from splunk_ao.search import RecordType, Search +from splunk_ao.shared.base import StateManagementMixin, SyncState +from splunk_ao.shared.exceptions import ValidationError +from splunk_ao.shared.project_resolver import _resolve_project +from splunk_ao.shared.query_result import QueryResult + +if TYPE_CHECKING: + from splunk_ao.shared.column import ColumnCollection -import warnings +logger = logging.getLogger(__name__) + +# Mapping from RecordType (plural) to RootType (singular) +RECORD_TYPE_TO_ROOT_TYPE = { + RecordType.SPAN: RootType.SPAN, + RecordType.TRACE: RootType.TRACE, + RecordType.SESSION: RootType.SESSION, +} -from splunk_ao.log_stream import LogStream __all__ = ["AgentStream"] -class AgentStream(LogStream): +class AgentStream(StateManagementMixin): """ - Object-centric interface for Splunk AO agent streams. + Object-centric interface for Galileo log streams. - ``AgentStream`` is the new name for what was previously called a - *Log Stream*. All functionality is identical; only the name has - changed. Use ``AgentStream`` for all new code. + This class provides an intuitive way to work with Galileo log streams, + offering methods for managing log streams and their associated metrics. - See ``splunk_ao.log_stream.LogStream`` for the full API reference — - every method, property, and class-method is inherited unchanged. + Attributes + ---------- + created_at (datetime.datetime): When the log stream was created. + created_by (str): The user who created the log stream. + id (str): The unique log stream identifier. + name (str): The log stream name. + project_id (str): The ID of the project this log stream belongs to. + project_name (str | None): The name of the project. May be None if the log stream + was retrieved using project_id only, as the API doesn't return this field. + updated_at (datetime.datetime): When the log stream was last updated. + additional_properties (dict): Additional properties of the log stream. Examples -------- - from splunk_ao import AgentStream + # Create a new log stream and persist it + log_stream = AgentStream(name="Production Logs", project_name="My AI Project").create() + + # Get an existing log stream + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") - # Create and persist a new agent stream - stream = AgentStream(name="prod-traces", project_name="my-project").create() + # AgentStreams can also be created through Project instances + from splunk_ao.project import Project - # Retrieve an existing stream - stream = AgentStream.get(name="prod-traces", project_name="my-project") + project = Project.get(name="My AI Project") + log_stream = project.create_log_stream(name="Production Logs") - # List streams for a project - streams = AgentStream.list(project_name="my-project") + # Enable metrics on the log stream + from splunk_ao.schema.metrics import SplunkAOMetrics + local_metrics = log_stream.enable_evaluators([ + SplunkAOMetrics.correctness, + SplunkAOMetrics.completeness, + "context_relevance" + ]) - # Enable evaluators on the stream - from splunk_ao import SplunkAOMetrics - stream.set_metrics([SplunkAOMetrics.correctness, SplunkAOMetrics.completeness]) + # Refresh log stream state from API + log_stream.refresh() """ + created_at: datetime | None + created_by: str | None + id: str | None + name: str + project_id: str | None + project_name: str | None # May not be available when retrieved from API + updated_at: datetime | None + additional_properties: dict[str, Any] # TODO: We need to validate if we will keep this one. + + def __str__(self) -> str: + """String representation of the log stream.""" + return f"AgentStream(name='{self.name}', id='{self.id}', project_id='{self.project_id}')" + + def __repr__(self) -> str: + """Detailed string representation of the log stream.""" + return f"AgentStream(name='{self.name}', id='{self.id}', project_id='{self.project_id}', created_at='{self.created_at}')" + + def __init__(self, name: str, *, project_id: str | None = None, project_name: str | None = None) -> None: + """ + Initialize a AgentStream instance locally. + + Creates a local log stream object that exists only in memory until .create() + is called to persist it to the API. + + Args: + name (str): The name of the log stream to create. + project_id (Optional[str]): The project ID. If neither project_id nor project_name is provided, + falls back to SPLUNK_AO_PROJECT_ID or SPLUNK_AO_PROJECT environment variables. + project_name (Optional[str]): The project name. If neither project_id nor project_name is provided, + falls back to SPLUNK_AO_PROJECT environment variable. + + Raises + ------ + ValidationError: If name is not provided. + ValueError: If project cannot be resolved at create() time (no explicit param and no env fallback). + + Examples + -------- + # Create by project ID + log_stream = AgentStream(name="Production Logs", project_id="project-123") + + # Create by project name + log_stream = AgentStream(name="Production Logs", project_name="My AI Project") + + # Create using SPLUNK_AO_PROJECT environment variable + log_stream = AgentStream(name="Production Logs") + """ + super().__init__() + + if not name: + raise ValidationError("'name' must be provided to create a log stream.") + + # Initialize attributes locally + self.name = name + self.project_id = project_id + self.project_name = project_name # May be None; not returned by API + self.id = None + self.created_at = None + self.created_by = None + self.updated_at = None + self.additional_properties = {} + + # Set initial state + self._set_state(SyncState.LOCAL_ONLY) + + def create(self) -> AgentStream: + """ + Persist this log stream to the API. + + Returns + ------- + AgentStream: This log stream instance with updated attributes from the API. + + Raises + ------ + NotFoundError: If the project cannot be found (no explicit param and no env fallback). + Exception: If the API call fails. + + Examples + -------- + log_stream = AgentStream(name="Production Logs", project_name="My AI Project").create() + assert log_stream.is_synced() + """ + if not self.name: + raise ValueError("Log stream name is not set. Cannot create log stream without a name.") + + # Note: project_id and project_name can both be None here — resolution happens below + # via _resolve_project, which reads SPLUNK_AO_PROJECT_ID / SPLUNK_AO_PROJECT env vars. + + try: + logger.info(f"AgentStream.create: name='{self.name}' project_id='{self.project_id}' - started") + + # _resolve_project raises NotFoundError when no project can be found, matching + # the contract of AgentStream.get() and AgentStream.list(). + project_obj = _resolve_project(self.project_id, self.project_name) + + # Update project info from resolved project + self.project_id = project_obj.id + if self.project_name is None: + self.project_name = project_obj.name + + agent_streams_svc = AgentStreams() + created_log_stream = agent_streams_svc.create( + name=self.name, + project_id=self.project_id, # always set by _resolve_project above + ) + + # Update attributes from response + self.created_at = created_log_stream.created_at + self.created_by = created_log_stream.created_by + self.id = created_log_stream.id + self.name = created_log_stream.name + self.project_id = created_log_stream.project_id + self.updated_at = created_log_stream.updated_at + self.additional_properties = created_log_stream.additional_properties + # Note: project_name is preserved if it was set, but API doesn't return it + + # Set state to synced + self._set_state(SyncState.SYNCED) + logger.info(f"AgentStream.create: id='{self.id}' - completed") + return self + except Exception as e: + self._set_state(SyncState.FAILED_SYNC, error=e) + logger.error(f"AgentStream.create: name='{self.name}' - failed: {e}") + raise + + @classmethod + def _create_empty(cls) -> AgentStream: + """Internal constructor bypassing __init__ for API hydration.""" + instance = cls.__new__(cls) + super(AgentStream, instance).__init__() + return instance + + @classmethod + def _from_api_response(cls, retrieved_log_stream: Any) -> AgentStream: + """ + Factory method to create a AgentStream instance from an API response. + + Args: + retrieved_log_stream: The log stream data retrieved from the API. + + Returns + ------- + AgentStream: A new AgentStream instance populated with the API data. + """ + instance = cls._create_empty() + instance.created_at = retrieved_log_stream.created_at + instance.created_by = retrieved_log_stream.created_by + instance.id = retrieved_log_stream.id + instance.name = retrieved_log_stream.name + instance.project_id = retrieved_log_stream.project_id + instance.updated_at = retrieved_log_stream.updated_at + instance.additional_properties = retrieved_log_stream.additional_properties + instance.project_name = None # API doesn't return project_name + # Set state to synced since we just retrieved from API + instance._set_state(SyncState.SYNCED) + return instance + + @classmethod + def get(cls, *, name: str, project_id: str | None = None, project_name: str | None = None) -> AgentStream | None: + """ + Get an existing log stream by name. + + Args: + name (str): The log stream name. + project_id (Optional[str]): The project ID. If neither project_id nor project_name is provided, + falls back to SPLUNK_AO_PROJECT_ID or SPLUNK_AO_PROJECT environment variables. + project_name (Optional[str]): The project name. If neither project_id nor project_name is provided, + falls back to SPLUNK_AO_PROJECT environment variable. + + Returns + ------- + Optional[AgentStream]: The log stream if found, None otherwise. + + Raises + ------ + NotFoundError: If the project cannot be found (no explicit param and no env fallback). + + Examples + -------- + # Get by project name + log_stream = AgentStream.get( + name="Production Logs", + project_name="My AI Project" + ) + + # Get by project ID + log_stream = AgentStream.get( + name="Production Logs", + project_id="project-123" + ) + + # Get using SPLUNK_AO_PROJECT environment variable + log_stream = AgentStream.get(name="Production Logs") + """ + project_obj = _resolve_project(project_id, project_name) + + agent_streams_svc = AgentStreams() + retrieved_log_stream = agent_streams_svc.get(name=name, project_id=project_obj.id) + if retrieved_log_stream is None: + return None + + instance = cls._from_api_response(retrieved_log_stream) + # Set project_name from resolved project + instance.project_name = project_obj.name + return instance + + @classmethod + def list( + cls, + *, + project_id: str | None = None, + project_name: str | None = None, + limit: Unset | int = 100, + starting_token: Unset | int = 0, + ) -> list[AgentStream]: + """ + List log streams for a project. -# Convenience: expose the deprecated ``LogStream`` name with a warning -def __getattr__(name: str) -> object: - if name == "LogStream": - warnings.warn( - "splunk_ao.agent_stream.LogStream is deprecated; " - "import AgentStream from splunk_ao.agent_stream instead.", - DeprecationWarning, - stacklevel=2, + Returns a single page of results. Use `starting_token` (from + `next_starting_token` on a prior response) to fetch subsequent pages. + + Args: + project_id (Optional[str]): The project ID. If neither project_id nor project_name is provided, + falls back to SPLUNK_AO_PROJECT_ID or SPLUNK_AO_PROJECT environment variables. + project_name (Optional[str]): The project name. If neither project_id nor project_name is provided, + falls back to SPLUNK_AO_PROJECT environment variable. + limit (Union[Unset, int]): Maximum number of log streams to return per page. Defaults to 100. + starting_token (Union[Unset, int]): Pagination token to start from. Defaults to 0 (first page). + + Returns + ------- + List[AgentStream]: A page of log streams for the project. + + Raises + ------ + NotFoundError: If the project cannot be found (no explicit param and no env fallback). + + Examples + -------- + # List by project name + log_streams = AgentStream.list(project_name="My AI Project") + + # List by project ID + log_streams = AgentStream.list(project_id="project-123") + + # List using SPLUNK_AO_PROJECT environment variable + log_streams = AgentStream.list() + + # Cap the number of returned log streams + log_streams = AgentStream.list(project_name="My AI Project", limit=3) + + # Fetch the next page + page_2 = AgentStream.list(project_name="My AI Project", starting_token=100) + """ + project_obj = _resolve_project(project_id, project_name) + + agent_streams_svc = AgentStreams() + retrieved_log_streams = agent_streams_svc.list( + project_id=project_obj.id, limit=limit, starting_token=starting_token ) - return LogStream - raise AttributeError(f"module 'splunk_ao.agent_stream' has no attribute {name!r}") + + instances = [cls._from_api_response(retrieved_log_stream) for retrieved_log_stream in retrieved_log_streams] + # Set project_name from resolved project for all instances + for instance in instances: + instance.project_name = project_obj.name + return instances + + def refresh(self) -> None: + """ + Refresh this log stream's state from the API. + + Updates all attributes with the latest values from the remote API + and sets the state to SYNCED. + + Raises + ------ + ValueError: If the log stream ID or project_id is not set. + Exception: If the API call fails or the log stream no longer exists. + + Examples + -------- + log_stream.refresh() + assert log_stream.is_synced() + """ + if self.id is None: + raise ValueError("Log stream ID is not set. Cannot refresh a local-only log stream.") + + if self.project_id is None: + raise ValueError("Project ID is not set. Cannot refresh log stream without project_id.") + + try: + logger.debug(f"AgentStream.refresh: id='{self.id}' - started") + agent_streams_svc = AgentStreams() + retrieved_log_stream = agent_streams_svc.get(id=self.id, project_id=self.project_id) + + if retrieved_log_stream is None: + raise ValueError(f"Log stream with id '{self.id}' no longer exists") + + # Update all attributes from response + self.created_at = retrieved_log_stream.created_at + self.created_by = retrieved_log_stream.created_by + self.id = retrieved_log_stream.id + self.name = retrieved_log_stream.name + self.project_id = retrieved_log_stream.project_id + self.updated_at = retrieved_log_stream.updated_at + self.additional_properties = retrieved_log_stream.additional_properties + # Note: project_name is preserved from before refresh since API doesn't return it + + # Set state to synced + self._set_state(SyncState.SYNCED) + logger.debug(f"AgentStream.refresh: id='{self.id}' - completed") + except Exception as e: + self._set_state(SyncState.FAILED_SYNC, error=e) + logger.error(f"AgentStream.refresh: id='{self.id}' - failed: {e}") + raise + + def get_metrics(self) -> builtins.list[str]: + """ + Get the list of metrics currently enabled on this log stream. + + Returns + ------- + list[str]: List of metric names currently enabled. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + log_stream = AgentStream.get(name="Production Logs", project_name="My Project") + current_metrics = log_stream.get_metrics() + print(f"Currently enabled: {current_metrics}") + """ + logger.info(f"AgentStream.get_metrics: id='{self.id}' - started") + config = SplunkAOConfig.get() + + settings = get_settings_projects_project_id_runs_run_id_scorer_settings_get.sync( + project_id=self.project_id, run_id=self.id, client=config.api_client + ) + + if settings is None or not hasattr(settings, "scorers"): + logger.info(f"AgentStream.get_metrics: id='{self.id}' - no metrics enabled") + return [] + + # Extract metric names from scorer configs + metric_names = [scorer.name for scorer in settings.scorers] + logger.info(f"AgentStream.get_metrics: id='{self.id}' found {len(metric_names)} metrics - completed") + return metric_names + + def set_metrics( + self, metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str] + ) -> builtins.list[LocalMetricConfig]: + """ + Set (replace) the metrics on this log stream. + + This replaces any existing metrics with the new list. Alias for enable_metrics + with clearer naming intent. + + Args: + metrics: List of metrics to set. Supports: + - SplunkAOMetrics enum values (e.g., SplunkAOMetrics.correctness) + - Metric objects (including from Metric.get(id="...")) + - LocalMetricConfig objects for custom scoring functions + - String names of built-in metrics + + Returns + ------- + List[LocalMetricConfig]: Local metric configurations that must be + computed client-side. + + Raises + ------ + ValueError: If any specified metrics are unknown. + + Examples + -------- + from splunk_ao import Metric, AgentStream + + log_stream = AgentStream.get(name="Production Logs", project_name="My Project") + + # Set metrics (replaces existing) + log_stream.set_metrics([ + Metric.metrics.correctness, + Metric.metrics.completeness, + Metric.get(id="metric-from-console-uuid"), # From console + ]) + """ + try: + logger.info(f"AgentStream.enable_evaluators: id='{self.id}' metrics={[str(m) for m in metrics]} - started") + agent_streams_svc = AgentStreams() + log_stream = agent_streams_svc.get(name=self.name, project_id=self.project_id) + if log_stream is None: + raise ValueError(f"Log stream '{self.name}' not found") + result = log_stream.enable_evaluators(metrics) + # Set state to synced after successful operation + self._set_state(SyncState.SYNCED) + logger.info(f"AgentStream.enable_evaluators: id='{self.id}' - completed") + return result + except Exception as e: + self._set_state(SyncState.FAILED_SYNC, error=e) + logger.error(f"AgentStream.enable_evaluators: id='{self.id}' - failed: {e}") + raise + + def query( + self, + record_type: RecordType, + filters: builtins.list[FilterType] | None = None, + sort: LogRecordsSortClause | None = None, + limit: int = 100, + starting_token: int = 0, + ) -> QueryResult: + """ + Query records in this log stream. + + This method provides a convenient way to search spans, traces, or sessions + within the current log stream without needing to specify project_id or log_stream_id. + + Args: + record_type: The type of records to query (SPAN, TRACE, or SESSION). + filters: A list of filters to apply to the query. + sort: A sort clause to order the query results. + limit: The maximum number of records to return. + starting_token: The token for the next page of results. + + Returns + ------- + QueryResult: A list-like object containing the query results with pagination support. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + from splunk_ao.search import RecordType + + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") + + # Query with column-based filters and sort + results = log_stream.query( + record_type=RecordType.SPAN, + filters=[ + log_stream.span_columns["input"].contains("largest"), + log_stream.span_columns["metrics/completeness_gpt"].greater_than(0.8), + log_stream.span_columns["created_at"].after("2024-01-01") + ], + sort=log_stream.span_columns["created_at"].descending(), + limit=50 + ) + + # Access results like a list + for record in results: + print(record["id"], record["input"]) + + # Get specific record + first_record = results[0] + + # Pagination + if results.has_next_page: + next_results = results.next_page() + """ + if self.id is None: + raise ValueError("Log stream ID is not set. Cannot query a local-only log stream.") + if self.project_id is None: + raise ValueError("Project ID is not set. Cannot query log stream without project_id.") + + logger.debug(f"AgentStream.query: id='{self.id}' record_type='{record_type.value}' limit={limit} - started") + + # Capture project_id and log_stream_id for use in pagination function + project_id = self.project_id + log_stream_id = self.id + + search_service = Search() + response = search_service.query( + project_id=project_id, + record_type=record_type, + log_stream_id=log_stream_id, + filters=filters, + sort=sort, + limit=limit, + starting_token=starting_token, + ) + + # Create a query function that returns raw response for pagination + def query_fn( + record_type: RecordType, + filters: builtins.list[FilterType] | None, + sort: LogRecordsSortClause | None, + limit: int, + starting_token: int, + ) -> Any: + return Search().query( + project_id=project_id, + record_type=record_type, + log_stream_id=log_stream_id, + filters=filters, + sort=sort, + limit=limit, + starting_token=starting_token, + ) + + # Wrap the response in QueryResult for easy access and pagination + return QueryResult(response=response, query_fn=query_fn, record_type=record_type, filters=filters, sort=sort) + + def get_spans( + self, + filters: builtins.list[FilterType] | None = None, + sort: LogRecordsSortClause | None = None, + limit: int = 100, + starting_token: int = 0, + ) -> QueryResult: + """ + Query spans in this log stream. + + This is a convenience method that queries for spans specifically. + + Args: + filters: A list of filters to apply to the query. + sort: A sort clause to order the query results. + limit: The maximum number of records to return. + starting_token: The token for the next page of results. + + Returns + ------- + QueryResult: A list-like object containing the span query results with pagination support. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") + + # Get spans with filters and sorting + spans = log_stream.get_spans( + filters=[ + log_stream.span_columns["input"].contains("world"), + log_stream.span_columns["metrics/num_input_tokens"].greater_than(10) + ], + sort=log_stream.span_columns["created_at"].descending(), + limit=50 + ) + + # Iterate over results + for span in spans: + print(span["id"], span["input"]) + + # Pagination + if spans.has_next_page: + more_spans = spans.next_page() + """ + logger.debug(f"AgentStream.get_spans: id='{self.id}' limit={limit} - started") + return self.query( + record_type=RecordType.SPAN, filters=filters, sort=sort, limit=limit, starting_token=starting_token + ) + + def get_traces( + self, + filters: builtins.list[FilterType] | None = None, + sort: LogRecordsSortClause | None = None, + limit: int = 100, + starting_token: int = 0, + ) -> QueryResult: + """ + Query traces in this log stream. + + This is a convenience method that queries for traces specifically. + + Args: + filters: A list of filters to apply to the query. + sort: A sort clause to order the query results. + limit: The maximum number of records to return. + starting_token: The token for the next page of results. + + Returns + ------- + QueryResult: A list-like object containing the trace query results with pagination support. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") + + # Get traces with filters + traces = log_stream.get_traces( + filters=[ + log_stream.trace_columns["input"].contains("largest"), + log_stream.trace_columns["created_at"].after("2024-01-01") + ], + sort=log_stream.trace_columns["created_at"].descending(), + limit=50 + ) + + # Access like a list + for trace in traces: + print(trace["id"], trace["input"]) + """ + logger.debug(f"AgentStream.get_traces: id='{self.id}' limit={limit} - started") + return self.query( + record_type=RecordType.TRACE, filters=filters, sort=sort, limit=limit, starting_token=starting_token + ) + + def get_sessions( + self, + filters: builtins.list[FilterType] | None = None, + sort: LogRecordsSortClause | None = None, + limit: int = 100, + starting_token: int = 0, + ) -> QueryResult: + """ + Query sessions in this log stream. + + This is a convenience method that queries for sessions specifically. + + Args: + filters: A list of filters to apply to the query. + sort: A sort clause to order the query results. + limit: The maximum number of records to return. + starting_token: The token for the next page of results. + + Returns + ------- + QueryResult: A list-like object containing the session query results with pagination support. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") + + # Get sessions with filters + sessions = log_stream.get_sessions( + filters=[ + log_stream.session_columns["model"].equals("gpt-4o-mini"), + log_stream.session_columns["metrics/num_traces"].greater_than(5) + ], + sort=log_stream.session_columns["created_at"].descending(), + limit=50 + ) + + # Work with results + for session in sessions: + print(session["id"], session["model"]) + """ + logger.debug(f"AgentStream.get_sessions: id='{self.id}' limit={limit} - started") + return self.query( + record_type=RecordType.SESSION, filters=filters, sort=sort, limit=limit, starting_token=starting_token + ) + + def export_records( + self, + record_type: RecordType = RecordType.TRACE, + filters: builtins.list[FilterType] | None = None, + sort: LogRecordsSortClause = LogRecordsSortClause(column_id="created_at", ascending=False), + export_format: LLMExportFormat = LLMExportFormat.JSONL, + column_ids: builtins.list[str] | None = None, + redact: bool = True, + ) -> Iterator[dict[str, Any]]: + """ + Export records from this log stream. + + This method provides a convenient way to export records without needing + to specify project_id or log_stream_id. + + Args: + record_type: The type of records to export (SPAN, TRACE, or SESSION). + filters: A list of filters to apply to the export. + sort: A sort clause to order the exported records. + export_format: The desired format for the exported data. + column_ids: A list of column IDs to include in the export. + redact: Redact sensitive data from the response. + + Returns + ------- + Iterator[dict[str, Any]]: An iterator that yields each record as a dictionary. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + from splunk_ao.search import RecordType + + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") + + # Export records with filters + for record in log_stream.export_records( + record_type=RecordType.SPAN, + filters=[ + log_stream.span_columns["model"].one_of(["gpt-4", "gpt-3.5-turbo", "gpt-4o-mini"]), + log_stream.span_columns["metrics/num_input_tokens"].greater_than(1) + ], + sort=log_stream.span_columns["created_at"].descending() + ): + print(record) + """ + if self.id is None: + raise ValueError("Log stream ID is not set. Cannot export from a local-only log stream.") + if self.project_id is None: + raise ValueError("Project ID is not set. Cannot export log stream without project_id.") + + # Convert RecordType to RootType for the export client + root_type = RECORD_TYPE_TO_ROOT_TYPE[record_type] + + logger.info( + f"AgentStream.export_records: id='{self.id}' record_type='{record_type.value}' " + f"export_format='{export_format.value}' - started" + ) + export_client = ExportClient() + return export_client.records( + project_id=self.project_id, + root_type=root_type, + filters=filters, + sort=sort, + export_format=export_format, + log_stream_id=self.id, + column_ids=column_ids, + redact=redact, + ) + + def context(self) -> Any: + """ + Get a galileo context manager for this log stream. + + This is a convenient method that returns a pre-configured splunk_ao_context + for this log stream, eliminating the need to specify project and log stream names. + + Returns + ------- + A context manager for Galileo logging configured with this log stream. + + Examples + -------- + log_stream = AgentStream.get( + name="Production Logs", + project_name="My AI Project" + ) + + with log_stream.context(): + # Your logging code here + response = openai_client.chat.completions.create(...) + """ + return splunk_ao_context(project=self.project.name if self.project else None, log_stream=self.name) + + def _get_columns(self, api_func: Any, error_msg: str) -> LogRecordsAvailableColumnsResponse: + """Helper method to retrieve available columns from the API.""" + if self.id is None: + raise ValueError("Log stream ID is not set. Cannot get columns from a local-only log stream.") + if self.project_id is None: + raise ValueError("Project ID is not set. Cannot get columns without project_id.") + + config = SplunkAOConfig.get() + body = LogRecordsAvailableColumnsRequest(log_stream_id=self.id) + response = api_func.sync(project_id=self.project_id, client=config.api_client, body=body) + if isinstance(response, HTTPValidationError): + raise response + if not response: + raise ValueError(error_msg) + return response + + @property + def project(self) -> Project | None: + """Get the project this log stream belongs to.""" + return Project.get(id=self.project_id) + + @property + def span_columns(self) -> ColumnCollection: + """ + Get available columns for spans in this log stream. + + Returns + ------- + ColumnCollection: A collection of columns available for spans, accessible by column ID. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") + columns = log_stream.span_columns + + # Access a specific column + input_column = columns["input"] + + # Filter using columns + spans = log_stream.get_spans( + filters=[columns["input"].contains("world")], + sort=columns["created_at"].descending() + ) + """ + response = self._get_columns( + spans_available_columns_projects_project_id_spans_available_columns_post, "Unable to retrieve span columns" + ) + columns = [Column(col) for col in response.columns] + return ColumnCollection(columns) + + @property + def session_columns(self) -> ColumnCollection: + """ + Get available columns for sessions in this log stream. + + Returns + ------- + ColumnCollection: A collection of columns available for sessions, accessible by column ID. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") + columns = log_stream.session_columns + + # Access a specific column + model_column = columns["model"] + + # Filter using columns + sessions = log_stream.get_sessions( + filters=[columns["model"].equals("gpt-4o-mini")], + sort=columns["created_at"].descending() + ) + """ + response = self._get_columns( + sessions_available_columns_projects_project_id_sessions_available_columns_post, + "Unable to retrieve session columns", + ) + columns = [Column(col) for col in response.columns] + return ColumnCollection(columns) + + @property + def trace_columns(self) -> ColumnCollection: + """ + Get available columns for traces in this log stream. + + Returns + ------- + ColumnCollection: A collection of columns available for traces, accessible by column ID. + + Raises + ------ + ValueError: If the log stream lacks required id or project_id attributes. + + Examples + -------- + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") + columns = log_stream.trace_columns + + # Access a specific column + input_column = columns["input"] + + # Filter using columns + traces = log_stream.get_traces( + filters=[columns["input"].contains("largest")], + sort=columns["created_at"].descending() + ) + """ + response = self._get_columns( + traces_available_columns_projects_project_id_traces_available_columns_post, + "Unable to retrieve trace columns", + ) + columns = [Column(col) for col in response.columns] + return ColumnCollection(columns) + + +# Import at end to avoid circular import (project.py imports AgentStream) +from splunk_ao.project import Project # noqa: E402 +from splunk_ao.resources.api.run_scorer_settings import ( # noqa: E402 + get_settings_projects_project_id_runs_run_id_scorer_settings_get, +) +from splunk_ao.shared.column import Column, ColumnCollection # noqa: E402 diff --git a/src/splunk_ao/agent_streams.py b/src/splunk_ao/agent_streams.py index 28484c7e..2a6ebb6e 100644 --- a/src/splunk_ao/agent_streams.py +++ b/src/splunk_ao/agent_streams.py @@ -1,84 +1,618 @@ -""" -Agent Streams service layer — the renamed successor to Log Streams (HYBIM-730). - -Provides ``AgentStreams`` (service class) and the module-level convenience -functions ``get_agent_stream``, ``list_agent_streams``, ``create_agent_stream``, -and ``enable_evaluators``. - -The old ``log_streams`` module and its symbols remain available but are -deprecated. -""" -from __future__ import annotations - import builtins -import warnings - -from splunk_ao.agent_stream import AgentStream -from splunk_ao.log_streams import LogStreams, create_log_stream, get_log_stream, list_log_streams +from typing import overload + +from splunk_ao.config import SplunkAOConfig +from splunk_ao.projects import Projects +from splunk_ao.resources.api.log_stream import ( + create_log_stream_projects_project_id_log_streams_post, + get_log_stream_projects_project_id_log_streams_log_stream_id_get, + list_log_streams_paginated_projects_project_id_log_streams_paginated_get, +) +from splunk_ao.resources.models.http_validation_error import HTTPValidationError +from splunk_ao.resources.models.log_stream_create_request import LogStreamCreateRequest +from splunk_ao.resources.models.log_stream_response import LogStreamResponse from splunk_ao.resources.types import Unset from splunk_ao.schema.metrics import LocalMetricConfig, Metric, SplunkAOMetrics +from splunk_ao.utils.env_helpers import _get_log_stream_from_env, _get_project_from_env +from splunk_ao.utils.log_config import get_logger +from splunk_ao.utils.metrics import create_metric_configs -__all__ = [ - "AgentStreams", - "create_agent_stream", - "enable_evaluators", - "get_agent_stream", - "list_agent_streams", -] +logger = get_logger(__name__) -class AgentStreams(LogStreams): +class AgentStream(LogStreamResponse): """ - Low-level service class for managing agent streams. + Log streams are used to organize logs within a project on the Galileo platform. + They provide a way to categorize and group related logs, making it easier to + analyze and monitor specific parts of your application or different environments + (e.g., production, staging, development). - ``AgentStreams`` is the new name for the ``LogStreams`` service class. - Inherits all methods from ``LogStreams`` and returns ``AgentStream`` - instances from listing/retrieval methods. + Attributes + ---------- + created_at : datetime.datetime + The timestamp when the log stream was created. + created_by : str + The identifier of the user who created the log stream. + id : str + The unique identifier of the log stream. + name : str + The name of the log stream. + project_id : str + The ID of the project this log stream belongs to. + updated_at : datetime.datetime + The timestamp when the log stream was last updated. + additional_properties : dict + Additional properties associated with the log stream. Examples -------- - from splunk_ao.agent_streams import AgentStreams + ```python + # Create a new log stream in a project + from splunk_ao.log_streams import create_agent_stream + + # Create by project ID + log_stream = create_agent_stream(name="Production Logs", project_id="project-123") + + # Create by project name + log_stream = create_agent_stream(name="Production Logs", project_name="My AI Project") + + # Get a log stream by name + from splunk_ao.log_streams import get_agent_stream + log_stream = get_agent_stream(name="Production Logs", project_name="My AI Project") + + # List all log streams in a project + from splunk_ao.log_streams import list_agent_streams + log_streams = list_agent_streams(project_name="My AI Project") + for stream in log_streams: + logger.info(f"Log Stream: {stream.name} (ID: {stream.id})") + + # Use a log stream with the context manager + from splunk_ao.openai import openai + from splunk_ao import splunk_ao_context + + with splunk_ao_context(project="My AI Project", log_stream="Production Logs"): + response = openai.chat.completions.create( + model="gpt-4o", + messages=[{"role": "user", "content": "Hello, world!"}] + ) - svc = AgentStreams() - stream = svc.get(name="prod-traces", project_name="my-project") - streams = svc.list(project_name="my-project") + # Enable metrics on a log stream - RECOMMENDED APPROACH + from splunk_ao.log_streams import enable_evaluators + from splunk_ao.schema.metrics import SplunkAOMetrics + + # Set environment variables first + # export SPLUNK_AO_LOG_STREAM="Production Logs" + # export SPLUNK_AO_PROJECT="My AI Project" + + # Clean and simple - just pass the metrics! + local_metrics = enable_evaluators([ + SplunkAOMetrics.correctness, + SplunkAOMetrics.completeness, + "context_relevance" + ]) + + # Alternative: Use explicit parameters + local_metrics = enable_evaluators( + agent_stream_name="Production Logs", + project_name="My AI Project", + metrics=["correctness", "completeness"] + ) + ``` """ + def __init__(self, log_stream: None | LogStreamResponse = None): + """ + Initialize a AgentStream instance. + + Parameters + ---------- + log_stream : Union[None, LogStreamResponse], optional + The log stream data to initialize from. If None, creates an empty log stream instance. + Defaults to None. + """ + if log_stream is not None: + super().__init__( + created_at=log_stream.created_at, + id=log_stream.id, + name=log_stream.name, + project_id=log_stream.project_id, + updated_at=log_stream.updated_at, + created_by=log_stream.created_by, + ) + self.additional_properties = log_stream.additional_properties.copy() + return + + def enable_evaluators( + self, metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str] + ) -> builtins.list[LocalMetricConfig]: + """ + Enable metrics directly on this log stream instance. + + This is the most intuitive and clean way to enable metrics when you already have a + AgentStream object. The method leverages the log stream's existing project_id and id + attributes, eliminating the need for redundant parameter specification and reducing + the potential for errors. + + This approach is ideal for object-oriented workflows where you're working with AgentStream + instances directly, and it provides the clearest semantic meaning: "enable these metrics + on this specific log stream." + + Parameters + ---------- + metrics : builtins.list[Union[SplunkAOMetrics, Metric, LocalMetricConfig, str]] + List of metrics to enable on this log stream. Supports multiple input formats: + + - **SplunkAOMetrics enum values**: Built-in metrics like `SplunkAOMetrics.correctness` + - **Metric objects**: Custom metrics with optional version specifications + - **LocalMetricConfig objects**: Client-side metrics with custom scoring functions + - **String names**: Built-in metric names like "correctness" or "toxicity" + + Returns + ------- + builtins.list[LocalMetricConfig] + List of local metric configurations that must be computed client-side. + Server-side metrics are automatically registered with Galileo and don't + need to be returned since users don't interact with them. + + Raises + ------ + ValueError + - If this AgentStream instance lacks required `id` or `project_id` attributes + - If any specified metrics are unknown or unavailable + - If there are issues with metric configuration or registration + GalileoHTTPException + If there are network or API errors when communicating with Galileo services + + Examples + -------- + Basic usage with built-in metrics: + + ```python + from splunk_ao.log_streams import AgentStreams + from splunk_ao.schema.metrics import SplunkAOMetrics + + # Get a log stream first + log_streams = AgentStreams() + log_stream = log_streams.get(name="Production Logs", project_name="My AI Project") + + # Enable metrics directly - clean and intuitive! + local_metrics = log_stream.enable_evaluators([ + SplunkAOMetrics.correctness, + SplunkAOMetrics.completeness, + "context_relevance", + "toxicity" + ]) + + logger.info(f"Server-side metrics enabled automatically") + logger.info(f"Need to process {len(local_metrics)} local metrics") + ``` + + Advanced usage with custom metrics: + + ```python + from splunk_ao.schema.metrics import Metric, LocalMetricConfig + + def custom_scorer(trace_or_span): + return 0.75 # Your scoring logic + + local_metrics = log_stream.enable_evaluators([ + SplunkAOMetrics.correctness, + "completeness", + Metric(name="domain_relevance", version=3), + LocalMetricConfig(name="custom_metric", scorer_fn=custom_scorer) + ]) + + # Process local metrics if any + for local_metric in local_metrics: + logger.info(f"Need to process local metric: {local_metric.name}") + ``` + + Notes + ----- + **Requirements:** + + - The AgentStream instance must have valid `id` and `project_id` attributes + - These are automatically set when retrieving AgentStream objects via AgentStreams methods + + **Recommended Usage:** + + - Use this method when you already have a AgentStream object + - More intuitive than specifying project/log stream names again + - Cleaner object-oriented design pattern + """ + if not hasattr(self, "id") or not hasattr(self, "project_id"): + raise ValueError("Log stream must have id and project_id to enable metrics") + + _, local_metrics = create_metric_configs(self.project_id, self.id, metrics) + return local_metrics + + +class AgentStreams: + config: SplunkAOConfig + + def __init__(self) -> None: + self.config = SplunkAOConfig.get() + + @overload + def list( + self, *, project_id: str, limit: Unset | int = 100, starting_token: Unset | int = 0 + ) -> builtins.list[AgentStream]: ... + + @overload + def list( + self, *, project_name: str, limit: Unset | int = 100, starting_token: Unset | int = 0 + ) -> builtins.list[AgentStream]: ... + + def list( + self, + *, + project_id: str | None = None, + project_name: str | None = None, + limit: Unset | int = 100, + starting_token: Unset | int = 0, + ) -> builtins.list[AgentStream]: + """ + Lists log streams. Exactly one of `project_id` or `project_name` must be provided. + + Returns a single page of results. Use `starting_token` (from + `next_starting_token` on a prior response) to fetch subsequent pages. + + Parameters + ---------- + project_id : Optional[str], optional + The ID of the project to list log streams for. + project_name : Optional[str], optional + The name of the project to list log streams for. + limit : Union[Unset, int], optional + The maximum number of log streams to return per page. Defaults to 100. + starting_token : Union[Unset, int], optional + The pagination token to start from. Defaults to 0 (first page). + + Returns + ------- + builtins.list[AgentStream] + A page of log streams. + + Raises + ------ + ValueError + If neither or both `project_id` and `project_name` are provided, + if the named project is not found, if the server returns a + validation error, or if the response is unexpectedly empty. + errors.UnexpectedStatus + If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. + httpx.TimeoutException + If the request takes longer than Client.timeout. + """ + if (project_id is None) == (project_name is None): + raise ValueError("Exactly one of 'project_id' or 'project_name' must be provided") + + if not project_id: + project = Projects().get(name=project_name) + if not project: + raise ValueError(f"Project {project_name} not found") + project_id = project.id + + response = list_log_streams_paginated_projects_project_id_log_streams_paginated_get.sync( + client=self.config.api_client, project_id=project_id, limit=limit, starting_token=starting_token + ) + + if isinstance(response, HTTPValidationError): + raise ValueError(f"Failed to list log streams: {response.detail}") + if response is None: + raise ValueError("Unexpected empty response while listing log streams") + + return [AgentStream(log_stream=log_stream) for log_stream in response.log_streams] + + # Page size used by `_list_all`. Larger than the default `list()` page size so + # full scans (name-based `get`, oldest-stream fallback) issue fewer round trips. + # The OpenAPI spec does not declare a server-side maximum, so this value is a + # heuristic. If the server ever rejects it (e.g. 422), `_list_all` now raises + # `ValueError` instead of silently truncating, so the regression is loud. + _LIST_ALL_PAGE_SIZE = 500 + + def _list_all(self, *, project_id: str) -> builtins.list[AgentStream]: + """Internal helper: paginate through every page and return all log streams. + + Used by callers that need a globally-complete view (e.g. name-based lookup, + oldest-stream fallback). Public callers should use `list()` with explicit + pagination instead. + + Raises ValueError on server validation failures or unexpected protocol + errors so mid-pagination failures aren't silently swallowed into a + truncated result. + """ + all_log_streams: builtins.list[AgentStream] = [] + starting_token: int = 0 + seen_tokens: set[int] = {starting_token} + while True: + response = list_log_streams_paginated_projects_project_id_log_streams_paginated_get.sync( + client=self.config.api_client, + project_id=project_id, + starting_token=starting_token, + limit=self._LIST_ALL_PAGE_SIZE, + ) + if isinstance(response, HTTPValidationError): + raise ValueError(f"Failed to list log streams: {response.detail}") + if response is None: + raise ValueError("Unexpected empty response while paginating log streams") + + all_log_streams.extend(AgentStream(log_stream=log_stream) for log_stream in response.log_streams) + + next_token = response.next_starting_token + if next_token is None or isinstance(next_token, Unset) or not response.paginated: + break + # Progress guard: stop if we've already seen this token. Catches both + # repeated and non-advancing tokens without assuming monotonic-integer + # ordering, so the loop stays safe if the server ever switches to + # opaque cursor tokens. + if not isinstance(next_token, int) or next_token in seen_tokens: + break + seen_tokens.add(next_token) + starting_token = next_token + + return all_log_streams + + @overload + def get(self, *, id: str, project_id: str | None = None, project_name: str | None = None) -> AgentStream | None: ... + @overload + def get(self, *, name: str, project_id: str | None = None, project_name: str | None = None) -> AgentStream | None: ... + def get( + self, + *, + id: str | None = None, + name: str | None = None, + project_id: str | None = None, + project_name: str | None = None, + ) -> AgentStream | None: + """ + Retrieves a log stream by id or name. + + Parameters + ---------- + id : Optional[str], optional + The id of the log stream. Defaults to None. + name : Optional[str], optional + The name of the log stream. Defaults to None. + project_id : Optional[str], optional + The ID of the project. Defaults to None. + project_name : Optional[str], optional + The name of the project. Defaults to None. + + Returns + ------- + Optional[AgentStream] + The log stream if found, None otherwise. + + Raises + ------ + ValueError + If neither or both `id` and `name` are provided, or if neither or both + `project_id` and `project_name` are provided. + errors.UnexpectedStatus + If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. + httpx.TimeoutException + If the request takes longer than Client.timeout. + """ + if (id is None) == (name is None): + raise ValueError("Exactly one of 'id' or 'name' must be provided") + + if (project_id is None) == (project_name is None): + raise ValueError("Exactly one of 'project_id' or 'project_name' must be provided") + + if not project_id: + project = Projects().get(name=project_name) + if not project: + raise ValueError(f"Project {project_name} not found") + project_id = project.id + + if id: + log_stream_response = get_log_stream_projects_project_id_log_streams_log_stream_id_get.sync( + project_id=project_id, log_stream_id=id, client=self.config.api_client + ) + if not log_stream_response: + return None + return AgentStream(log_stream=log_stream_response) + + if name: + for log_stream in self._list_all(project_id=project_id): + if log_stream.name == name: + return log_stream + return None + + @overload + def create(self, name: str, *, project_id: str | None = None) -> AgentStream: ... + @overload + def create(self, name: str, *, project_name: str) -> AgentStream: ... + + def create(self, name: str, *, project_id: str | None = None, project_name: str | None = None) -> AgentStream: + """ + Creates a new log stream. Exactly one of `project_id` or `project_name` must be provided. + + Parameters + ---------- + name : str + The name of the log stream. + project_id : Optional[str], optional + The ID of the project to create the log stream in. Defaults to None. + project_name : Optional[str], optional + The name of the project to create the log stream in. Defaults to None. + + Returns + ------- + AgentStream + The created log stream. + + Raises + ------ + ValueError + If neither or both `project_id` and `project_name` are provided, or if the project is not found. + HTTPValidationError + If the server validation fails. + errors.UnexpectedStatus + If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. + httpx.TimeoutException + If the request takes longer than Client.timeout. + """ + if (project_id is None) == (project_name is None): + raise ValueError("Exactly one of 'project_id' or 'project_name' must be provided") + + if not project_id: + project = Projects().get(name=project_name) + if not project: + raise ValueError(f"Project {project_name} not found") + project_id = project.id + + body = LogStreamCreateRequest(name=name) + response = create_log_stream_projects_project_id_log_streams_post.sync( + project_id=project_id, client=self.config.api_client, body=body + ) + + if isinstance(response, HTTPValidationError): + raise response + + if not response: + raise ValueError("Unable to create log stream") + + return AgentStream(log_stream=response) + + def enable_evaluators( + self, + *, + agent_stream_name: str | None = None, + project_name: str | None = None, + metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str], + ) -> builtins.list[LocalMetricConfig]: + """ + Enable metrics for a log stream by configuring scorers. + + The project name can be provided via the 'project_name' parameter or the + SPLUNK_AO_PROJECT environment variable. + + The log stream name can be provided via the 'agent_stream_name' parameter or the + SPLUNK_AO_LOG_STREAM environment variable. + + Parameters + ---------- + agent_stream_name : Optional[str], optional + The name of the log stream. Takes precedence over the SPLUNK_AO_LOG_STREAM environment variable. Defaults to None. + project_name : Optional[str], optional + The name of the project. Takes precedence over the SPLUNK_AO_PROJECT environment variable. Defaults to None. + metrics : builtins.list[Union[SplunkAOMetrics, Metric, LocalMetricConfig, str]] + List of metrics to enable. Can include: + - SplunkAOMetrics enum values (e.g., SplunkAOMetrics.correctness) + - Metric objects with name and optional version + - LocalMetricConfig objects for custom local metrics + - String names of built-in metrics + + Returns + ------- + tuple[builtins.list[ScorerConfig], builtins.list[LocalMetricConfig]] + A tuple containing the configured scorer configs and local metric configs. + + Raises + ------ + ValueError + If log stream or project cannot be found, or if metrics are unknown. + + Examples + -------- + ```python + # Enable built-in metrics with explicit parameters + from splunk_ao.log_streams import AgentStreams + from splunk_ao.schema.metrics import SplunkAOMetrics + + log_streams = AgentStreams() + scorer_configs, local_metrics = log_streams.enable_evaluators( + agent_stream_name="Production Logs", + project_name="My AI Project", + metrics=[ + SplunkAOMetrics.correctness, + SplunkAOMetrics.completeness, + "context_relevance", + ], + ) + + # Enable metrics using environment variables + # export SPLUNK_AO_LOG_STREAM="Production Logs" + # export SPLUNK_AO_PROJECT="My AI Project" + scorer_configs, local_metrics = log_streams.enable_evaluators( + metrics=["correctness", "completeness"] + ) + + # Enable custom metrics with mixed parameters + from splunk_ao.schema.metrics import Metric, LocalMetricConfig + + def custom_scorer(trace_or_span): + return 0.85 # Custom scoring logic + + # export SPLUNK_AO_PROJECT="My AI Project" + scorer_configs, local_metrics = log_streams.enable_evaluators( + agent_stream_name="Production Logs", # Explicit log stream + # project_name from env var + metrics=[ + Metric(name="my_custom_metric", version=2), + LocalMetricConfig(name="local_scorer", scorer_fn=custom_scorer) + ] + ) + ``` + """ + # Apply environment variable fallbacks + project_name = project_name or _get_project_from_env() + agent_stream_name = agent_stream_name or _get_log_stream_from_env() + + # Get project using environment fallbacks + project_obj = Projects().get_with_env_fallbacks(name=project_name) + if not project_obj: + raise ValueError(f"Project '{project_name}' not found") + + # Get log stream - error out if not found + if not agent_stream_name: + raise ValueError("agent_stream_name must be provided (or set SPLUNK_AO_LOG_STREAM env var)") + log_stream = self.get(name=agent_stream_name, project_name=project_obj.name) + if not log_stream: + raise ValueError(f"Log stream '{agent_stream_name}' not found in project '{project_obj.name}'") -# --------------------------------------------------------------------------- -# Module-level convenience functions -# --------------------------------------------------------------------------- + # Use the shared utility function directly + _, local_metrics = create_metric_configs(project_obj.id, log_stream.id, metrics) + return local_metrics + + +# +# Convenience methods +# def get_agent_stream( - *, - name: str | None = None, - project_id: str | None = None, - project_name: str | None = None, + *, name: str | None = None, project_id: str | None = None, project_name: str | None = None ) -> AgentStream | None: """ - Retrieve an agent stream by name. + Retrieves a log stream by name. Exactly one of `project_id` or `project_name` must be provided. Parameters ---------- - name: - The agent stream name. - project_id: - The project ID (mutually exclusive with *project_name*). - project_name: - The project name (mutually exclusive with *project_id*). + name : Optional[str], optional + The name of the log stream. Defaults to None. + project_id : Optional[str], optional + The ID of the project. Defaults to None. + project_name : Optional[str], optional + The name of the project. Defaults to None. Returns ------- - AgentStream | None - The agent stream if found, ``None`` otherwise. + Optional[AgentStream] + The log stream if found, None otherwise. + + Raises + ------ + ValueError + If neither or both `project_id` and `project_name` are provided. + errors.UnexpectedStatus + If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. + httpx.TimeoutException + If the request takes longer than Client.timeout. """ - result = get_log_stream(name=name, project_id=project_id, project_name=project_name) - if result is None: - return None - stream = AgentStream.__new__(AgentStream) - stream.__dict__.update(result.__dict__) - return stream + return AgentStreams().get(name=name, project_id=project_id, project_name=project_name) # type: ignore[arg-type] def list_agent_streams( @@ -89,64 +623,63 @@ def list_agent_streams( starting_token: Unset | int = 0, ) -> builtins.list[AgentStream]: """ - List agent streams for a project. + Lists log streams. Exactly one of `project_id` or `project_name` must be provided. + + Returns a single page of results. Use `starting_token` (from + `next_starting_token` on a prior response) to fetch subsequent pages. Parameters ---------- - project_id: - The project ID (mutually exclusive with *project_name*). - project_name: - The project name (mutually exclusive with *project_id*). - limit: - Maximum number of results per page. Defaults to 100. - starting_token: - Pagination token. Defaults to 0. + project_id : str + The id of the project. + project_name : str + The name of the project. + limit : Union[Unset, int], optional + The maximum number of log streams to return per page. Defaults to 100. + starting_token : Union[Unset, int], optional + The pagination token to start from. Defaults to 0 (first page). Returns ------- - list[AgentStream] - A page of agent streams. + builtins.list[AgentStream] + A page of log streams. + + Raises + ------ + errors.UnexpectedStatus + If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. + httpx.TimeoutException + If the request takes longer than Client.timeout. + """ - log_streams = list_log_streams( - project_id=project_id, - project_name=project_name, - limit=limit, - starting_token=starting_token, + return AgentStreams().list( # type: ignore[call-overload] + project_id=project_id, project_name=project_name, limit=limit, starting_token=starting_token ) - result: builtins.list[AgentStream] = [] - for ls in log_streams: - s = AgentStream.__new__(AgentStream) - s.__dict__.update(ls.__dict__) - result.append(s) - return result -def create_agent_stream( - name: str, - project_id: str | None = None, - project_name: str | None = None, -) -> AgentStream: +def create_agent_stream(name: str, project_id: str | None = None, project_name: str | None = None) -> AgentStream: """ - Create a new agent stream. + Creates a new log stream. Exactly one of `project_id` or `project_name` must be provided. Parameters ---------- - name: - The agent stream name. - project_id: - The project ID (mutually exclusive with *project_name*). - project_name: - The project name (mutually exclusive with *project_id*). + name : str + The name of the log stream. Returns ------- AgentStream - The created agent stream. + The created project. + + Raises + ------ + errors.UnexpectedStatus + If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. + httpx.TimeoutException + If the request takes longer than Client.timeout. + """ - ls = create_log_stream(name=name, project_id=project_id, project_name=project_name) - stream = AgentStream.__new__(AgentStream) - stream.__dict__.update(ls.__dict__) - return stream + return AgentStreams().create(name=name, project_id=project_id, project_name=project_name) # type: ignore[call-overload] def enable_evaluators( @@ -156,53 +689,104 @@ def enable_evaluators( metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str], ) -> builtins.list[LocalMetricConfig]: """ - Enable evaluators (formerly *metrics*) on an agent stream. + Enable metrics for a log stream with flexible parameter and environment variable support. + + This unified function supports both explicit parameters and environment variable fallbacks, + making it perfect for all use cases - from production CI/CD pipelines to development testing. - Falls back to the ``SPLUNK_AO_LOG_STREAM`` and ``SPLUNK_AO_PROJECT`` - environment variables when *agent_stream_name* / *project_name* are - not provided explicitly. + **Flexible Usage Patterns:** + - **Environment-only**: Just pass metrics, names from env vars (production/CI) + - **Explicit parameters**: Specify project/log stream names directly (development) + - **Mixed approach**: Combine explicit params with environment fallbacks + + Environment Variables (Optional Fallbacks) + ------------------------------------------ + SPLUNK_AO_PROJECT : str + The name of the Galileo project (used when project_name not provided) + SPLUNK_AO_LOG_STREAM : str + The name of the log stream (used when agent_stream_name not provided) Parameters ---------- - agent_stream_name: - The agent stream name. Falls back to ``SPLUNK_AO_LOG_STREAM`` env var. - project_name: - The project name. Falls back to ``SPLUNK_AO_PROJECT`` env var. - metrics: - Evaluators to enable. Accepts ``SplunkAOMetrics`` enum values, - ``Metric`` objects, ``LocalMetricConfig`` objects, or string names. + agent_stream_name : Optional[str], optional + The name of the log stream. Takes precedence over SPLUNK_AO_LOG_STREAM environment variable. + If None, will use SPLUNK_AO_LOG_STREAM env var. Defaults to None. + project_name : Optional[str], optional + The name of the project. Takes precedence over SPLUNK_AO_PROJECT environment variable. + If None, will use SPLUNK_AO_PROJECT env var. Defaults to None. + metrics : builtins.list[Union[SplunkAOMetrics, Metric, LocalMetricConfig, str]] + List of metrics to enable on the log stream. Can include: + - SplunkAOMetrics enum values (e.g., SplunkAOMetrics.correctness) + - Metric objects with name and optional version for custom metrics + - LocalMetricConfig objects for client-side custom scoring functions + - String names of built-in metrics (e.g., "correctness", "toxicity") Returns ------- - list[LocalMetricConfig] - Local evaluator configurations that must be computed client-side. - """ - from splunk_ao.log_streams import LogStreams + builtins.list[LocalMetricConfig] + List of local metric configurations that must be computed client-side. + Server-side metrics are automatically registered with Galileo and don't + need to be returned since users don't interact with them. + + Raises + ------ + ValueError + If log stream or project cannot be found, or if any specified metrics are unknown. + errors.UnexpectedStatus + If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. + httpx.TimeoutException + If the request takes longer than Client.timeout. - return LogStreams().enable_metrics( - log_stream_name=agent_stream_name, - project_name=project_name, - metrics=metrics, + Examples + -------- + ```python + # Enable built-in metrics with explicit parameters + from splunk_ao.log_streams import enable_evaluators + from splunk_ao.schema.metrics import SplunkAOMetrics + + local_metrics = enable_evaluators( + agent_stream_name="Production Logs", + project_name="My AI Project", + metrics=[ + SplunkAOMetrics.correctness, + SplunkAOMetrics.completeness, + "context_relevance", + ], ) + # Enable metrics using environment variables only + # export SPLUNK_AO_LOG_STREAM="Production Logs" + # export SPLUNK_AO_PROJECT="My AI Project" + local_metrics = enable_evaluators(metrics=["correctness", "completeness"]) + + # Enable custom and local metrics with environment variable fallbacks + from splunk_ao.schema.metrics import Metric, LocalMetricConfig + from galileo_core.schemas.logging.step import StepType + + def response_length_scorer(trace_or_span): + '''Custom metric that scores based on response length''' + if hasattr(trace_or_span, "output") and trace_or_span.output: + return min(len(trace_or_span.output) / 100.0, 1.0) # Normalize 0-1 + return 0.0 + + local_metrics = enable_evaluators( + agent_stream_name="Development Logs", + metrics=[ + SplunkAOMetrics.correctness, + "toxicity", + Metric(name="my_custom_metric", version=2), + LocalMetricConfig( + name="response_length", + scorer_fn=response_length_scorer, + scorable_types=[StepType.llm], + aggregatable_types=[StepType.trace], + ), + ], + ) -# --------------------------------------------------------------------------- -# Deprecated aliases for the old ``log_streams`` convenience functions -# --------------------------------------------------------------------------- - -def __getattr__(name: str) -> object: - _deprecated = { - "get_log_stream": ("get_agent_stream", get_agent_stream), - "list_log_streams": ("list_agent_streams", list_agent_streams), - "create_log_stream": ("create_agent_stream", create_agent_stream), - "enable_metrics": ("enable_evaluators", enable_evaluators), - } - if name in _deprecated: - new_name, obj = _deprecated[name] - warnings.warn( - f"splunk_ao.agent_streams.{name} is deprecated; use {new_name} instead.", - DeprecationWarning, - stacklevel=2, - ) - return obj - raise AttributeError(f"module 'splunk_ao.agent_streams' has no attribute {name!r}") + # Process local metrics + for local_metric in local_metrics: + logger.info(f"Need to process local metric: {local_metric.name}") + ``` + """ + return AgentStreams().enable_evaluators(agent_stream_name=agent_stream_name, project_name=project_name, metrics=metrics) diff --git a/src/splunk_ao/evaluator.py b/src/splunk_ao/evaluator.py index b9b3e360..034ae5fb 100644 --- a/src/splunk_ao/evaluator.py +++ b/src/splunk_ao/evaluator.py @@ -1,181 +1,1317 @@ -""" -Evaluators — the renamed successor to Metrics (HYBIM-730). - -``Evaluator`` and its concrete subclasses (``LlmEvaluator``, ``CodeEvaluator``, -``LocalEvaluator``, ``SplunkAOEvaluator``) are the canonical names going -forward. The old ``Metric``-prefixed names are kept as deprecated aliases in -``splunk_ao.__init__`` and will be removed in a future major release. - -The underlying implementation lives in ``splunk_ao.metric``; the API endpoints -still use the ``/scorers`` path (server-side rename is tracked separately). -""" from __future__ import annotations -import warnings +import builtins +import json +import logging +import os +import time +from abc import ABC +from collections.abc import Callable +from datetime import datetime +from typing import TYPE_CHECKING, Any -from splunk_ao.metric import ( - BuiltInMetrics, - CodeMetric, - LlmMetric, - LocalMetric, - Metric, - SplunkAOMetric, +if TYPE_CHECKING: + from splunk_ao.model import Model + +from galileo_core.schemas.logging.span import Span +from galileo_core.schemas.logging.step import StepType +from galileo_core.schemas.logging.trace import Trace +from galileo_core.schemas.shared.metric import MetricValueType +from splunk_ao.config import SplunkAOConfig +from splunk_ao.configuration import Configuration +from splunk_ao.evaluators import Evaluators +from splunk_ao.resources.api.data import ( + create_code_scorer_version_scorers_scorer_id_version_code_post, + create_scorers_post, + get_validate_code_scorer_task_result_scorers_code_validate_task_id_get, + update_scorers_scorer_id_patch, + validate_code_scorer_scorers_code_validate_post, +) +from splunk_ao.resources.models import ( + BodyCreateCodeScorerVersionScorersScorerIdVersionCodePost, + BodyValidateCodeScorerScorersCodeValidatePost, + CreateScorerRequest, + HTTPValidationError, + OutputTypeEnum, + ScorerTypes, + TaskResultStatus, + UpdateScorerRequest, ) +from splunk_ao.resources.models.invalid_result import InvalidResult +from splunk_ao.resources.types import UNSET, File, Unset +from splunk_ao.schema.metrics import LocalMetricConfig, SplunkAOMetrics +from splunk_ao.schema.metrics import Metric as SchemaMetric +from splunk_ao.scorers import Scorers +from splunk_ao.shared.base import StateManagementMixin, SyncState +from splunk_ao.shared.exceptions import APIError, ValidationError -__all__ = [ - "BuiltInEvaluators", - "CodeEvaluator", - "Evaluator", - "LlmEvaluator", - "LocalEvaluator", - "SplunkAOEvaluator", -] +logger = logging.getLogger(__name__) +# Code validation polling parameters are configurable via Configuration: +# - Configuration.code_validation_timeout (env: SPLUNK_AO_CODE_VALIDATION_TIMEOUT) - default: 60.0s +# - Configuration.code_validation_initial_delay (env: SPLUNK_AO_CODE_VALIDATION_INITIAL_DELAY) - default: 5.0s +# - Configuration.code_validation_max_delay (env: SPLUNK_AO_CODE_VALIDATION_MAX_DELAY) - default: 30.0s +# - Configuration.code_validation_backoff_multiplier (env: SPLUNK_AO_CODE_VALIDATION_BACKOFF_MULTIPLIER) - default: 1.5 -class BuiltInEvaluators(BuiltInMetrics): - """ - Provides attribute-style access to built-in Splunk AO evaluators. - This is the renamed successor to ``BuiltInMetrics``. Access built-in - evaluators via ``Evaluator.evaluators``. +class BuiltInEvaluators: + """ + Provides convenient access to built-in Galileo metrics (formerly "scorers"). Examples -------- - from splunk_ao import Evaluator - Evaluator.evaluators.correctness - Evaluator.evaluators.completeness + from splunk_ao.metric import Evaluator + + # Access built-in metrics + Evaluator.metrics.correctness + Evaluator.metrics.completeness + Evaluator.metrics.toxicity """ + def __getattr__(self, name: str) -> SplunkAOMetrics: + """Allow attribute-style access to built-in metrics.""" + # Try to find the metric by name (enum names match UI-visible names) + for scorer in SplunkAOMetrics: + if scorer.name == name: + return scorer + raise AttributeError(f"Built-in metric '{name}' not found. Available: {[s.name for s in SplunkAOMetrics]}") -class Evaluator(Metric): + def __dir__(self) -> list[str]: + """Return list of available metric names for autocomplete.""" + return [scorer.name for scorer in SplunkAOMetrics] + + +# Backwards-compatible alias +BuiltInScorers = BuiltInEvaluators + + +class Evaluator(StateManagementMixin, ABC): """ - Base class for all Splunk AO evaluators. + Base class for all Galileo metrics. - ``Evaluator`` is the new name for what was previously called a *Metric*. - All functionality is inherited from ``splunk_ao.metric.Metric`` unchanged. + This is an abstract base class that defines common attributes and methods + for all metric types. Use one of the concrete metric classes instead: - Use one of the concrete subclasses for new code: + - **SplunkAOEvaluator**: Built-in Galileo scorers (access via Evaluator.scorers) + - **LlmEvaluator**: Custom LLM-based metrics with prompt templates + - **LocalEvaluator**: Local function-based metrics + - **CodeEvaluator**: Code-based metrics (future support) - - ``SplunkAOEvaluator`` — built-in Splunk AO scorers (``Evaluator.evaluators.*``) - - ``LlmEvaluator`` — custom LLM-judge evaluators - - ``LocalEvaluator`` — local function-based evaluators - - ``CodeEvaluator`` — code-based evaluators + Common Attributes + ----------------- + id (str | None): The unique metric identifier (UUID). + name (str): The metric name. + scorer_type (ScorerTypes | None): The type of scorer. + description (str): Description of the metric. + tags (list[str]): Tags associated with the metric. + created_at (datetime | None): When the metric was created. + updated_at (datetime | None): When the metric was last updated. + version (int | None): Evaluator version number. Class Attributes ---------------- - evaluators : BuiltInEvaluators - Access built-in Splunk AO evaluators. + metrics (BuiltInEvaluators): Access built-in Galileo metrics. Examples -------- - from splunk_ao import Evaluator, AgentStream, SplunkAOMetrics + # 1. Use built-in Galileo scorers + from splunk_ao import Evaluator, SplunkAOEvaluator, LlmEvaluator, LocalEvaluator, LogStream - stream = AgentStream.get(name="prod-traces", project_name="my-project") - stream.set_metrics([ - Evaluator.evaluators.correctness, - Evaluator.evaluators.completeness, + log_stream = LogStream.get(name="my-stream", project_name="my-project") + log_stream.set_metrics([ + Evaluator.metrics.correctness, + Evaluator.metrics.completeness, ]) - # Get an existing evaluator by name - ev = Evaluator.get(name="factuality-checker") + # 2. Create custom LLM metric + llm_metric = LlmEvaluator( + name="response_quality", + prompt="Rate the quality...", + model="gpt-4o-mini", + judges=3, + ).create() + + # 3. Create local function-based metric + def my_scorer(trace_or_span): + return 0.5 - # List all evaluators - evaluators = Evaluator.list() + local_metric = LocalEvaluator( + name="response_length", + scorer_fn=my_scorer, + ) """ - evaluators = BuiltInEvaluators() + # Class attribute for built-in metrics (preferred name) + metrics = BuiltInEvaluators() + + # Backwards-compatible property for legacy name + scorers = metrics + + # Type annotations for common instance attributes + id: str | None + name: str + scorer_type: ScorerTypes | None + description: str + tags: list[str] + created_at: datetime | None + updated_at: datetime | None + version: int | None + + # Scorer defaults - available for LLM and built-in Galileo metrics + # These are returned by the API in the ScorerDefaults object + model: str | None + judges: int | None + cot_enabled: bool | None + + def __init__( + self, name: str, *, description: str = "", tags: list[str] | None = None, version: int | None = None + ) -> None: + """ + Initialize a base Evaluator instance with common attributes. + + Args: + name: The name of the metric. + description: Description of the metric. + tags: Tags associated with the metric. + version: Specific version to reference (for existing metrics). + """ + super().__init__() + self.name = name + self.description = description + self.tags = tags if tags is not None else [] + self.version = version + self.id = None + self.created_at = None + self.updated_at = None + self.scorer_type = None + + # Initialize scorer defaults (populated from API for LLM and Galileo metrics) + self.model = None + self.judges = None + self.cot_enabled = None + + self._set_state(SyncState.LOCAL_ONLY) + + @staticmethod + def _parse_output_type( + output_type: str | OutputTypeEnum | None, default: OutputTypeEnum | None = None + ) -> OutputTypeEnum | None: + """Map a string or OutputTypeEnum to OutputTypeEnum, with an optional fallback default.""" + if output_type is None: + return default + if isinstance(output_type, OutputTypeEnum): + return output_type + _map = { + "percentage": OutputTypeEnum.PERCENTAGE, + "boolean": OutputTypeEnum.BOOLEAN, + "categorical": OutputTypeEnum.CATEGORICAL, + "count": OutputTypeEnum.COUNT, + "discrete": OutputTypeEnum.DISCRETE, + "freeform": OutputTypeEnum.FREEFORM, + "multilabel": OutputTypeEnum.MULTILABEL, + } + return _map.get(output_type.lower(), default) + + @classmethod + def _create_metric_from_type(cls, scorer_type: ScorerTypes) -> Evaluator: + """ + Create the appropriate Evaluator subclass instance based on scorer_type. + + This is a factory method that centralizes the logic for instantiating + the correct metric subclass based on the scorer type returned from the API. + + Args: + scorer_type: The scorer type from the API response. + + Returns + ------- + Evaluator: An uninitialized instance of the appropriate subclass + (LlmEvaluator, CodeEvaluator, or SplunkAOEvaluator). + + Examples + -------- + instance = Evaluator._create_metric_from_type(ScorerTypes.LLM) + # Returns: LlmEvaluator instance + """ + if scorer_type == ScorerTypes.LLM: + return LlmEvaluator.__new__(LlmEvaluator) + if scorer_type == ScorerTypes.CODE: + return CodeEvaluator.__new__(CodeEvaluator) + # Default to SplunkAOEvaluator for built-in scorers (LUNA, PRESET, etc.) + return SplunkAOEvaluator.__new__(SplunkAOEvaluator) + + @classmethod + def get(cls, *, id: str | None = None, name: str | None = None) -> Evaluator | None: + """ + Get an existing metric by ID or name. + + Returns the appropriate subclass instance based on scorer_type. + + Args: + id: The metric ID (UUID). + name: The metric name. + + Returns + ------- + Optional[Evaluator]: The metric if found (SplunkAOEvaluator, LlmEvaluator, or CodeEvaluator), None otherwise. + + Raises + ------ + ValidationError: If neither or both id and name are provided. + + Examples + -------- + # Get by name - returns appropriate subclass + metric = Evaluator.get(name="factuality-checker") + + # Get by ID + metric = Evaluator.get(id="abc-123-def") + """ + if id is not None and name is not None: + raise ValidationError("Cannot specify both id and name") + if id is None and name is None: + raise ValidationError("Must specify either id or name") + + scorers_service = Scorers() + + if name is not None: + scorers = scorers_service.list(name=name) + if not scorers: + return None + retrieved_scorer = next((s for s in scorers if s.name == name), None) + if retrieved_scorer is None: + return None + else: + assert id is not None + scorers = scorers_service.list() + retrieved_scorer = next((s for s in scorers if s.id == id), None) + if retrieved_scorer is None: + return None + + # Create appropriate subclass instance based on scorer_type + instance = cls._create_metric_from_type(retrieved_scorer.scorer_type) + StateManagementMixin.__init__(instance) + instance._populate_from_scorer_response(retrieved_scorer) + instance._set_state(SyncState.SYNCED) + return instance + + @classmethod + def list( + cls, *, name_filter: str | None = None, scorer_types: list[ScorerTypes] | None = None + ) -> builtins.list[Evaluator]: + """ + List metrics with optional filtering. + + Returns appropriate subclass instances based on scorer_type. + + Args: + name_filter: Filter metrics by exact name match. + scorer_types: Filter by scorer types. + + Returns + ------- + list[Evaluator]: List of metrics matching the criteria (with appropriate subclass types). + + Examples + -------- + # List all metrics + metrics = Evaluator.list() + + # List LLM metrics only + metrics = Evaluator.list(scorer_types=[ScorerTypes.LLM]) + + # List by name + metrics = Evaluator.list(name_filter="factuality") + """ + logger.debug(f"Evaluator.list: name_filter='{name_filter}' types={scorer_types} - started") + scorers_service = Scorers() + retrieved_scorers = scorers_service.list(name=name_filter, types=scorer_types) + logger.debug(f"Evaluator.list: found {len(retrieved_scorers)} metrics - completed") + + result: builtins.list[Evaluator] = [] + for retrieved_scorer in retrieved_scorers: + # Create appropriate subclass instance based on scorer_type + instance = cls._create_metric_from_type(retrieved_scorer.scorer_type) + StateManagementMixin.__init__(instance) + instance._populate_from_scorer_response(retrieved_scorer) + instance._set_state(SyncState.SYNCED) + result.append(instance) + + return result + + @classmethod + def delete_by_name(cls, name: str) -> None: + """ + Delete a metric by name without retrieving it first. + + This is more efficient than calling `Evaluator.get(name=...).delete()` + when you only need to delete and don't need the metric object. - # ``metrics`` class attribute is inherited from Metric and intentionally - # left as-is. The new canonical accessor is ``Evaluator.evaluators``. + Args: + name: The name of the metric to delete. + Raises + ------ + ValueError: If no metric with the given name exists. -class LlmEvaluator(LlmMetric): + Examples + -------- + # Delete without retrieving first + Evaluator.delete_by_name("old-metric") + + # Alternative (less efficient) + metric = Evaluator.get(name="old-metric") + metric.delete() + """ + logger.info(f"Evaluator.delete_by_name: name='{name}' - started") + try: + metrics_service = Evaluators() + metrics_service.delete_evaluator(name=name) + logger.info(f"Evaluator.delete_by_name: name='{name}' - completed") + except Exception as e: + logger.error(f"Evaluator.delete_by_name: name='{name}' - failed: {e}") + raise + + def _populate_from_scorer_response(self, scorer_response: Any) -> None: + """Populate instance attributes from a ScorerResponse object.""" + # Pre-compute optional common attributes + description = ( + "" + if isinstance(scorer_response.description, Unset) or scorer_response.description is None + else scorer_response.description + ) + created_at = None if isinstance(scorer_response.created_at, Unset) else scorer_response.created_at + updated_at = None if isinstance(scorer_response.updated_at, Unset) else scorer_response.updated_at + + # Extract defaults - available for LLM and built-in Galileo metrics + # These are returned by the API for preset scorers too + if not isinstance(scorer_response.defaults, Unset) and scorer_response.defaults is not None: + model = scorer_response.defaults.model_name if hasattr(scorer_response.defaults, "model_name") else None + judges = scorer_response.defaults.num_judges if hasattr(scorer_response.defaults, "num_judges") else None + cot_enabled = ( + scorer_response.defaults.cot_enabled if hasattr(scorer_response.defaults, "cot_enabled") else None + ) + else: + model = None + judges = None + cot_enabled = None + + self._sync_attrs( + id=scorer_response.id, + name=scorer_response.name, + scorer_type=scorer_response.scorer_type, + tags=scorer_response.tags, + version=None, + description=description, + created_at=created_at, + updated_at=updated_at, + model=model, + judges=judges, + cot_enabled=cot_enabled, + ) + + # LLM-specific attributes (only set if this is an LlmEvaluator) + if isinstance(self, LlmEvaluator): + output_type = None if isinstance(scorer_response.output_type, Unset) else scorer_response.output_type + prompt = None if isinstance(scorer_response.user_prompt, Unset) else scorer_response.user_prompt + + # Extract scoreable node types + if not isinstance(scorer_response.scoreable_node_types, Unset) and scorer_response.scoreable_node_types: + try: + node_level = StepType(scorer_response.scoreable_node_types[0]) + except (ValueError, IndexError): + node_level = None + else: + node_level = None + + ground_truth = ( + False + if isinstance(scorer_response.ground_truth, Unset) or scorer_response.ground_truth is None + else scorer_response.ground_truth + ) + + self._sync_attrs(output_type=output_type, prompt=prompt, node_level=node_level, ground_truth=ground_truth) + + # Code-specific attributes (only set if this is a CodeEvaluator) + if isinstance(self, CodeEvaluator): + # Extract scoreable node types + if not isinstance(scorer_response.scoreable_node_types, Unset) and scorer_response.scoreable_node_types: + try: + code_node_level = StepType(scorer_response.scoreable_node_types[0]) + except (ValueError, IndexError): + code_node_level = None + else: + code_node_level = None + + code_output_type = None if isinstance(scorer_response.output_type, Unset) else scorer_response.output_type + + self._sync_attrs(node_level=code_node_level, output_type=code_output_type) + + def update(self, **kwargs: Any) -> Evaluator: + """ + Update this metric's properties on the API. + + Only ``name``, ``description``, and ``tags`` can be updated via this method. + On success the instance is updated with the API response and returned in SYNCED state. + + Parameters + ---------- + **kwargs : Any + Fields to update. Supported keys: ``name``, ``description``, ``tags``. + + Returns + ------- + Evaluator: This metric instance with updated attributes from the API. + + Raises + ------ + ValidationError: If this is a local metric. + ValueError: If the metric ID is not set, the metric is deleted, + or the metric is in FAILED_SYNC state. + ValueError: If any unsupported fields are passed. + APIError: If the API returns a validation error or empty response. + Exception: If the API call fails (state is set to FAILED_SYNC). + + Examples + -------- + metric = Evaluator.get(name="factuality-checker") + metric.update(name="new-name", description="Updated description") + assert metric.is_synced() + """ + if isinstance(self, LocalEvaluator): + raise ValidationError("Local metrics don't exist on the server and can't be updated.") + if self.id is None: + raise ValueError("Evaluator ID is not set. Cannot update a local-only metric.") + if self.sync_state == SyncState.DELETED: + raise ValueError("Cannot update a deleted metric.") + if self.sync_state == SyncState.FAILED_SYNC: + raise ValueError( + "Cannot update a metric in FAILED_SYNC state. " + "Call refresh() to re-sync from the API, then retry your changes." + ) + + valid_fields = {"name", "description", "tags"} + invalid_fields = set(kwargs) - valid_fields + if invalid_fields: + raise ValueError(f"Invalid update fields: {sorted(invalid_fields)!r}. Valid fields: {sorted(valid_fields)}") + + body = UpdateScorerRequest( + name=kwargs.get("name", UNSET), description=kwargs.get("description", UNSET), tags=kwargs.get("tags", UNSET) + ) + + logger.info(f"Evaluator.update: id='{self.id}' name='{self.name}' - started") + try: + config = SplunkAOConfig.get() + response = update_scorers_scorer_id_patch.sync(scorer_id=self.id, client=config.api_client, body=body) + except Exception as e: + self._set_state(SyncState.FAILED_SYNC, error=e) + logger.error(f"Evaluator.update: id='{self.id}' - failed: {e}") + raise + + if isinstance(response, HTTPValidationError): + raise APIError(f"Evaluator update validation error: {response.detail}") + if response is None: + raise APIError(f"Evaluator update returned empty response for id '{self.id}'") + + self._populate_from_scorer_response(response) + self._set_state(SyncState.SYNCED) + logger.info(f"Evaluator.update: id='{self.id}' - completed") + return self + + def delete(self) -> None: + """ + Delete this metric. + + Only works for server-side metrics. Local metrics don't need deletion. + + Raises + ------ + ValidationError: If this is a local metric. + ValueError: If the metric is not synced. + + Examples + -------- + metric = Evaluator.get(name="factuality-checker") + metric.delete() + """ + if isinstance(self, LocalEvaluator): + raise ValidationError("Local metrics don't exist on the server and can't be deleted.") + + if self.id is None: + raise ValueError("Evaluator ID is not set. Cannot delete a local-only metric.") + + try: + logger.info(f"Evaluator.delete: id='{self.id}' name='{self.name}' - started") + metrics_service = Evaluators() + metrics_service.delete_evaluator(name=self.name) + self._set_state(SyncState.DELETED) + logger.info(f"Evaluator.delete: id='{self.id}' - completed") + except Exception as e: + self._set_state(SyncState.FAILED_SYNC, error=e) + logger.error(f"Evaluator.delete: id='{self.id}' - failed: {e}") + raise + + def refresh(self) -> None: + """ + Refresh this metric's state from the API. + + Updates all attributes with the latest values from the remote API. + + Raises + ------ + ValidationError: If this is a local metric. + ValueError: If the metric is not synced. + Exception: If the API call fails or the metric no longer exists. + + Examples + -------- + metric.refresh() + assert metric.is_synced() + """ + if isinstance(self, LocalEvaluator): + raise ValidationError("Local metrics don't exist on the server and can't be refreshed.") + + if self.id is None: + raise ValueError("Evaluator ID is not set. Cannot refresh a local-only metric.") + + try: + logger.debug(f"Evaluator.refresh: id='{self.id}' - started") + scorers_service = Scorers() + scorers = scorers_service.list() + retrieved_scorer = next((s for s in scorers if s.id == self.id), None) + + if retrieved_scorer is None: + raise ValueError(f"Evaluator with id '{self.id}' no longer exists") + + self._populate_from_scorer_response(retrieved_scorer) + self._set_state(SyncState.SYNCED) + logger.debug(f"Evaluator.refresh: id='{self.id}' - completed") + except Exception as e: + self._set_state(SyncState.FAILED_SYNC, error=e) + logger.error(f"Evaluator.refresh: id='{self.id}' - failed: {e}") + raise + + def to_legacy_metric(self) -> SchemaMetric: + """ + Convert to legacy splunk_ao.schema.metrics.Evaluator format. + + This enables backward compatibility with existing code that uses + the legacy Evaluator class. + + Returns + ------- + SchemaMetric: Legacy metric object with name and version. + + Examples + -------- + metric = Evaluator.get(name="my-metric") + legacy = metric.to_legacy_metric() + # Use with existing APIs + """ + return SchemaMetric(name=self.name, version=self.version) + + def __str__(self) -> str: + """String representation of the metric.""" + type_name = self.__class__.__name__ + scorer_type_str = self.scorer_type.value if self.scorer_type else "unknown" + return f"{type_name}(name='{self.name}', id='{self.id}', scorer_type='{scorer_type_str}')" + + def __repr__(self) -> str: + """Detailed string representation of the metric.""" + type_name = self.__class__.__name__ + return f"{type_name}(name='{self.name}', id='{self.id}')" + + +# ============================================================================ +# Concrete Evaluator Types +# ============================================================================ + + +class LlmEvaluator(Evaluator): """ - LLM-based evaluator — the renamed successor to ``LlmMetric``. + LLM-based metric with custom prompt templates. + + This metric type allows you to create custom metrics evaluated by an LLM + judge using a prompt template. + + Attributes + ---------- + prompt (str | None): Prompt template for the LLM scorer. + model (str | None): Model name/alias to use for scoring (stored as string). + judges (int | None): Number of judges to use for scoring. + cot_enabled (bool | None): Whether chain-of-thought is enabled. + node_level (StepType | None): Node level for the metric. + output_type (OutputTypeEnum | None): Output type for the metric. - See ``splunk_ao.metric.LlmMetric`` for the full API reference. + Configuration + ------------- + Default values for `model` and `judges` can be configured via: + - Configuration.default_scorer_model (env: SPLUNK_AO_DEFAULT_SCORER_MODEL) + - Configuration.default_scorer_judges (env: SPLUNK_AO_DEFAULT_SCORER_JUDGES) Examples -------- - from splunk_ao import LlmEvaluator + # Create custom LLM metric with string model name + metric = LlmEvaluator( + name="response_quality", + prompt=''' + Rate the quality of this response on a scale of 1-10. + + Question: {input} + Answer: {output} - ev = LlmEvaluator( + Return only the numerical score (1-10). + ''', + model="gpt-4o-mini", # String model name + judges=3, + node_level=StepType.llm, + description="Rates response quality", + tags=["quality", "custom"], + output_type=OutputTypeEnum.PERCENTAGE, + cot_enabled=True, + ).create() + + # Or use a Model object from Integration + from splunk_ao.integration import Integration + gpt_model = Integration.openai.get_model(alias="gpt-4o-mini") + metric = LlmEvaluator( name="response_quality", - prompt="Rate the quality 1-10: {input} -> {output}", - model="gpt-4o-mini", + prompt="Rate quality 1-10: {input} -> {output}", + model=gpt_model, # Model object judges=3, ).create() """ + # Type annotations for LLM-specific attributes + prompt: str | None + model: str | None + judges: int | None + cot_enabled: bool | None + node_level: StepType | None + output_type: OutputTypeEnum | None + ground_truth: bool -class CodeEvaluator(CodeMetric): - """ - Code-based evaluator — the renamed successor to ``CodeMetric``. + def __init__( + self, + name: str, + *, + # LLM metric parameters (improved API) + prompt: str | None = None, + model: Model | str | None = None, + judges: int | None = None, + # Backward compatibility aliases + user_prompt: str | None = None, + model_name: str | None = None, + num_judges: int | None = None, + # LLM-specific parameters + node_level: StepType | None = None, + cot_enabled: bool | None = None, + output_type: str | OutputTypeEnum | None = None, + ground_truth: bool = False, + # Common parameters + description: str = "", + tags: list[str] | None = None, + version: int | None = None, + ) -> None: + """ + Initialize an LLM metric. + + Args: + name: The name of the metric. + prompt: Prompt template for LLM scorers (preferred over user_prompt). + model: Model object or model name string to use (preferred over model_name). + Defaults to Configuration.default_scorer_model. + judges: Number of judges (preferred over num_judges). Defaults to Configuration.default_scorer_judges. + user_prompt: [Deprecated] Use 'prompt' instead. + model_name: [Deprecated] Use 'model' instead. + num_judges: [Deprecated] Use 'judges' instead. + node_level: Node level for the metric. Defaults to StepType.llm. + cot_enabled: Whether chain-of-thought is enabled. Defaults to True. + output_type: Output type ("percentage", "boolean", etc.). + ground_truth: Whether the scorer requires ground truth (``reference_output``) from the dataset. + When True, the judge LLM receives the row's ground-truth value in its prompt. + description: Description of the metric. + tags: Tags associated with the metric. + version: Specific version to reference (for existing metrics). + + Raises + ------ + ValidationError: If prompt is not provided. + """ + super().__init__(name=name, description=description, tags=tags, version=version) + + # Handle parameter aliases (new names preferred) + final_prompt = prompt or user_prompt + + # Handle model parameter - extract alias from Model object if needed + if model is not None: + # Local import to avoid circular dependency + from splunk_ao.model import Model + + final_model = model.alias if isinstance(model, Model) else model + else: + final_model = model_name or Configuration.default_scorer_model + + final_judges = ( + judges + if judges is not None + else (num_judges if num_judges is not None else Configuration.default_scorer_judges) + ) + + if final_prompt is None: + raise ValidationError("'prompt' (or 'user_prompt') must be provided for LLM-based metrics.") + + # Initialize LLM-specific attributes + self.prompt = final_prompt + self.model = final_model # Now always a string (alias) + self.judges = final_judges + self.node_level = node_level or StepType.llm + self.cot_enabled = cot_enabled if cot_enabled is not None else True + + # Handle output_type (accept string or enum) + if isinstance(output_type, str): + self.output_type = Evaluator._parse_output_type(output_type, default=OutputTypeEnum.PERCENTAGE) + else: + self.output_type = output_type or OutputTypeEnum.BOOLEAN + + self.ground_truth = ground_truth + + self.scorer_type = ScorerTypes.LLM + + def create(self) -> LlmEvaluator: + """ + Persist this LLM metric to the API. + + Returns + ------- + LlmEvaluator: This metric instance with updated attributes from the API. + + Raises + ------ + ValidationError: If configuration is invalid. + Exception: If the API call fails. + + Examples + -------- + metric = LlmEvaluator( + name="quality_check", + prompt="Rate the quality...", + model="gpt-4o-mini" + ).create() + assert metric.is_synced() + """ + try: + logger.info(f"LlmEvaluator.create: name='{self.name}' - started") + + metrics_service = Evaluators() + created_version = metrics_service.create_custom_llm_evaluator( + name=self.name, + user_prompt=self.prompt or "", + node_level=self.node_level if self.node_level is not None else StepType.llm, + cot_enabled=self.cot_enabled if self.cot_enabled is not None else True, + model_name=self.model if self.model is not None else Configuration.default_scorer_model, + num_judges=self.judges if self.judges is not None else Configuration.default_scorer_judges, + description=self.description, + tags=self.tags, + output_type=self.output_type + if isinstance(self.output_type, OutputTypeEnum) + else OutputTypeEnum.BOOLEAN, + ground_truth=self.ground_truth, + ) + + # Update attributes from response without triggering dirty-tracking + self._sync_attrs( + id=str(created_version.scorer_id), + created_at=created_version.created_at, + updated_at=created_version.updated_at, + ) + + # Refresh to get full scorer details + self.refresh() - See ``splunk_ao.metric.CodeMetric`` for the full API reference. + logger.info(f"LlmEvaluator.create: id='{self.id}' - completed") + return self + except ValidationError: + raise + except Exception as e: + self._set_state(SyncState.FAILED_SYNC, error=e) + logger.error(f"LlmEvaluator.create: name='{self.name}' - failed: {e}") + raise + + def __repr__(self) -> str: + """Detailed string representation of the metric.""" + return f"LlmEvaluator(name='{self.name}', id='{self.id}', model='{self.model}', judges={self.judges})" + + +class CodeEvaluator(Evaluator): + r""" + Code-based metric. + + This metric type is for code-based scorers that execute custom code + to evaluate traces/spans. + + Attributes + ---------- + node_level (StepType | None): Node level for the metric. + code (str | None): The Python code for the scorer. + output_type (OutputTypeEnum | None): Output type for the metric. Examples -------- - from splunk_ao import CodeEvaluator + # Get existing code metric + metric = Evaluator.get(name="my-code-metric") + assert isinstance(metric, CodeEvaluator) - ev = CodeEvaluator( - name="custom_scorer", - code="def scorer_fn(step): return 1.0", + # Create code metric with inline code + metric = CodeEvaluator( + name="custom_code_scorer", + code="def scorer_fn(step_object):\\n return 1.0", + description="Custom code-based scorer", + tags=["custom", "code"], + node_level=StepType.llm, + output_type=OutputTypeEnum.PERCENTAGE, ).create() + + # Load code from file + metric = CodeEvaluator( + name="custom_code_scorer", + node_level=StepType.llm, + ).load_code("./scorers/my_scorer.py").create() """ + # Type annotations for code-specific attributes + node_level: StepType | None + code: str | None + output_type: OutputTypeEnum | None + required_metrics: list[str] | None + + def __init__( + self, + name: str, + *, + code: str | None = None, + node_level: StepType | None = None, + output_type: str | OutputTypeEnum | None = None, + required_metrics: list[str] | None = None, + description: str = "", + tags: list[str] | None = None, + version: int | None = None, + ) -> None: + """ + Initialize a Code metric. + + Args: + name: The name of the metric. + code: The Python code for the scorer (optional, can be set later or loaded from file). + node_level: Node level for the metric. Defaults to StepType.llm. + output_type: Output type for the metric ("percentage", "boolean", "categorical", + "count", "discrete"). Accepts string or OutputTypeEnum. + required_metrics: List of metric names that this code metric depends on. + description: Description of the metric. + tags: Tags associated with the metric. + version: Specific version to reference (for existing metrics). + """ + super().__init__(name=name, description=description, tags=tags, version=version) + + self.code = code + self.node_level = node_level or StepType.llm + self.required_metrics = required_metrics + self.scorer_type = ScorerTypes.CODE + + self.output_type = Evaluator._parse_output_type(output_type) + + def load_code(self, code_file_path: str) -> CodeEvaluator: + """ + Load code from a file into this metric instance. + + Args: + code_file_path: Path to the Python file containing the scorer code. + + Returns + ------- + CodeEvaluator: This metric instance with code loaded from the file (for chaining). + + Raises + ------ + ValidationError: If the code file doesn't exist or can't be read. + + Examples + -------- + # Load code from file + metric = CodeEvaluator( + name="custom_code_scorer", + node_level=StepType.llm, + ).load_code("./scorers/my_scorer.py").create() + """ + if not os.path.isfile(code_file_path): + raise ValidationError(f"Code file not found: {code_file_path}") + + try: + with open(code_file_path, encoding="utf-8") as f: + self.code = f.read() + except Exception as e: + raise ValidationError(f"Failed to read code file {code_file_path}: {e}") + + return self + + def _validate_code(self, config: SplunkAOConfig) -> str: + """ + Validate the code by submitting it to the validation endpoint and polling for results. + + Args: + config: The Galileo configuration with API client. + + Returns + ------- + str: The validation result as a JSON string to pass to create_code_scorer_version. + + Raises + ------ + ValidationError: If validation fails or the code is invalid. + ValueError: If the API returns an unexpected response. + """ + assert self.code is not None + assert self.node_level is not None + + # Step 1: Submit the code for validation + code_bytes = self.code.encode("utf-8") + code_file = File(payload=code_bytes, file_name="scorer.py") + validate_body = BodyValidateCodeScorerScorersCodeValidatePost( + file=code_file, scoreable_node_types=[self.node_level.value], required_scorers=self.required_metrics + ) + + validate_response = validate_code_scorer_scorers_code_validate_post.sync( + client=config.api_client, body=validate_body + ) + + if validate_response is None: + logger.debug("CodeEvaluator._validate_code: No response from validate_code_scorer") + raise ValueError("Failed to validate code: No response from API") + + task_id = validate_response.task_id + logger.debug(f"CodeEvaluator._validate_code: task_id='{task_id}' - validation started") + + # Step 2: Poll for validation result with time-based timeout + start_time_seconds = time.time() + attempt = 0 + while True: + elapsed_seconds = time.time() - start_time_seconds + timeout_seconds = Configuration.code_validation_timeout + if elapsed_seconds >= timeout_seconds: + raise ValidationError(f"Code validation timed out after {timeout_seconds:.0f} seconds") + + task_result = get_validate_code_scorer_task_result_scorers_code_validate_task_id_get.sync( + task_id=task_id, client=config.api_client + ) + + if task_result is None: + logger.debug(f"CodeEvaluator._validate_code: No response for task_id='{task_id}'") + raise ValueError("Failed to get validation result: No response from API") + + if task_result.status == TaskResultStatus.COMPLETED: + logger.debug(f"CodeEvaluator._validate_code: task_id='{task_id}' - validation completed") + + # Extract and validate the result + result = task_result.result + + # Handle string result (already serialized) + if isinstance(result, str): + return result + + # Handle ValidateRegisteredScorerResult or similar objects with to_dict + if hasattr(result, "to_dict"): + # Check if it's an invalid result (has error_message in nested result) + if hasattr(result, "result") and isinstance(result.result, InvalidResult): + raise ValidationError(f"Code validation failed: {result.result.error_message}") + # Return the result as JSON string + return json.dumps(result.to_dict()) + + raise ValueError(f"Unexpected validation result type: {type(result)}") + + if task_result.status == TaskResultStatus.FAILED: + error_msg = "Code validation failed" + if isinstance(task_result.result, str): + error_msg = f"Code validation failed: {task_result.result}" + raise ValidationError(error_msg) + + if task_result.status == TaskResultStatus.PENDING: + # Calculate delay with exponential backoff + delay_seconds = min( + Configuration.code_validation_initial_delay + * (Configuration.code_validation_backoff_multiplier**attempt), + Configuration.code_validation_max_delay, + ) + + logger.debug( + f"CodeEvaluator._validate_code: task_id='{task_id}' - pending " + f"(elapsed: {elapsed_seconds:.1f}s/{timeout_seconds:.0f}s, next delay: {delay_seconds:.2f}s)" + ) + time.sleep(delay_seconds) + attempt += 1 + else: + raise ValueError(f"Unknown task status: {task_result.status}") + + def create(self) -> CodeEvaluator: + r""" + Persist this Code metric to the API. + + This method validates the code first by submitting it to the validation + endpoint, polling for the result, and then creating the scorer with the + validated result. + + Returns + ------- + CodeEvaluator: This metric instance with updated attributes from the API. + + Raises + ------ + ValidationError: If code is not set, validation fails, or configuration is invalid. + Exception: If the API call fails. + + Examples + -------- + # Create with inline code + metric = CodeEvaluator( + name="custom_code_scorer", + code="def scorer_fn(step_object):\\n return 1.0", + node_level=StepType.llm, + ).create() + assert metric.is_synced() + + # Create by loading from file + metric = CodeEvaluator( + name="custom_code_scorer", + node_level=StepType.llm, + ).load_code("./scorers/my_scorer.py").create() + assert metric.is_synced() + """ + # Validate that code is set + if self.code is None: + raise ValidationError( + "Code is not set. Either pass 'code' to __init__() or use CodeEvaluator.load_code() to load from a file." + ) + + try: + logger.info(f"CodeEvaluator.create: name='{self.name}' - started") + + config = SplunkAOConfig.get() + + # Ensure node_level is set (should always be set in __init__, but checking for type safety) + assert self.node_level is not None + + # Step 1: Validate the code and get validation result + logger.debug(f"CodeEvaluator.create: name='{self.name}' - validating code") + validation_result = self._validate_code(config) + logger.debug(f"CodeEvaluator.create: name='{self.name}' - code validated successfully") + + # Step 2: Create the scorer + scorer_request = CreateScorerRequest( + name=self.name, + scorer_type=ScorerTypes.CODE, + description=self.description, + tags=self.tags, + scoreable_node_types=[self.node_level.value], + output_type=self.output_type, + required_scorers=self.required_metrics, + ) + + scorer_response = create_scorers_post.sync(client=config.api_client, body=scorer_request) + + if scorer_response is None: + logger.debug("CodeEvaluator.create: No response from create_scorers_post") + raise ValueError("Failed to create code-based metric: No response from API") -class LocalEvaluator(LocalMetric): + # Step 3: Create the code scorer version with file upload and validation result + # Convert the code string to bytes for file upload + code_bytes = self.code.encode("utf-8") + code_file = File(payload=code_bytes, file_name="scorer.py") + version_body = BodyCreateCodeScorerVersionScorersScorerIdVersionCodePost( + file=code_file, validation_result=validation_result + ) + + created_version = create_code_scorer_version_scorers_scorer_id_version_code_post.sync( + scorer_id=scorer_response.id, client=config.api_client, body=version_body + ) + + if created_version is None: + logger.debug( + "CodeEvaluator.create: No response from create_code_scorer_version_scorers_scorer_id_version_code_post" + ) + raise ValueError("Failed to create code-based metric: No response from API") + + # Update attributes from response + self.id = str(scorer_response.id) + self.created_at = scorer_response.created_at + self.updated_at = scorer_response.updated_at + + # Refresh to get full scorer details + self.refresh() + + logger.info(f"CodeEvaluator.create: id='{self.id}' - completed") + return self + except ValidationError: + raise + except Exception as e: + self._set_state(SyncState.FAILED_SYNC, error=e) + logger.error(f"CodeEvaluator.create: name='{self.name}' - failed: {e}") + raise + + def __repr__(self) -> str: + """Detailed string representation of the metric.""" + return f"CodeEvaluator(name='{self.name}', id='{self.id}'')" + + +class SplunkAOEvaluator(Evaluator): """ - Local function-based evaluator — the renamed successor to ``LocalMetric``. + Built-in Galileo scorer metric. - See ``splunk_ao.metric.LocalMetric`` for the full API reference. + This metric type represents Galileo's built-in scorers like correctness, + completeness, toxicity, etc. Access these via `Evaluator.metrics`. Examples -------- - from splunk_ao import LocalEvaluator + # Access built-in scorers + from splunk_ao import Evaluator, LogStream - def my_fn(trace): - return 0.5 + log_stream = LogStream.get(name="my-stream", project_name="my-project") + log_stream.set_metrics([ + Evaluator.metrics.correctness, + Evaluator.metrics.completeness, + Evaluator.metrics.toxicity, + ]) - ev = LocalEvaluator(name="my_scorer", scorer_fn=my_fn) + # Or get by name + metric = Evaluator.get(name="correctness") + assert isinstance(metric, SplunkAOEvaluator) """ + def __init__( + self, name: str, *, description: str = "", tags: list[str] | None = None, version: int | None = None + ) -> None: + """ + Initialize a Galileo metric. -class SplunkAOEvaluator(SplunkAOMetric): + Args: + name: The name of the metric. + description: Description of the metric. + tags: Tags associated with the metric. + version: Specific version to reference (for existing metrics). + """ + super().__init__(name=name, description=description, tags=tags, version=version) + # Galileo metrics can have various scorer types, set during population + + +class LocalEvaluator(Evaluator): """ - Built-in Splunk AO scorer evaluator — the renamed successor to ``SplunkAOMetric``. + Local function-based metric. + + This metric type uses a Python function to score traces/spans locally + without making API calls. Useful for simple, deterministic metrics. - Access built-in evaluators via ``Evaluator.evaluators.*``. + Attributes + ---------- + scorer_fn (Callable): Scoring function that takes a Trace or Span and returns either a + score, or a ``(score, metadata)`` tuple — when a tuple is returned, the metadata dict + is attached to the step under ``{name}_metadata`` for explainability. + scorable_types (list[StepType]): Types that can be scored. + aggregatable_types (list[StepType]): Types that can be aggregated. Examples -------- - from splunk_ao import Evaluator + # Create local function-based metric + def response_length_scorer(trace_or_span): + if hasattr(trace_or_span, "output") and trace_or_span.output: + return min(len(trace_or_span.output) / 100.0, 1.0) + return 0.0 - ev = Evaluator.get(name="correctness") - assert isinstance(ev, SplunkAOEvaluator) + local_metric = LocalEvaluator( + name="response_length", + scorer_fn=response_length_scorer, + scorable_types=[StepType.llm], + aggregatable_types=[StepType.trace], + ) + + # Or return (score, metadata) for explainability + EXPECTED = ["relevance", "accuracy", "completeness"] + def keyword_coverage(trace_or_span): + text = getattr(trace_or_span, "output", "") or "" + matched = [k for k in EXPECTED if k in text] + return len(matched) / len(EXPECTED), { + "matched": matched, + "missing": [k for k in EXPECTED if k not in text], + } + + # Use with log stream + log_stream.set_metrics([local_metric]) """ + # Type annotations for local metric attributes + scorer_fn: Callable[[Trace | Span], MetricValueType | tuple[MetricValueType, dict[str, Any]]] + scorable_types: list[StepType] + aggregatable_types: list[StepType] + + def __init__( + self, + name: str, + *, + scorer_fn: Callable[[Trace | Span], MetricValueType | tuple[MetricValueType, dict[str, Any]]], + scorable_types: list[StepType] | None = None, + aggregatable_types: list[StepType] | None = None, + description: str = "", + tags: list[str] | None = None, + ) -> None: + """ + Initialize a local function-based metric. + + Args: + name: The name of the metric. + scorer_fn: Scoring function for the metric. May return either a bare score, or a + ``(score, metadata)`` tuple where ``metadata`` is a JSON-serializable dict + surfaced under ``{name}_metadata`` on the step. + scorable_types: Step types that can be scored. Defaults to [StepType.llm]. + aggregatable_types: Step types for aggregation. Defaults to [StepType.trace]. + description: Description of the metric. + tags: Tags associated with the metric. + + Raises + ------ + ValidationError: If scorer_fn is not provided. + """ + super().__init__(name=name, description=description, tags=tags) -# --------------------------------------------------------------------------- -# Deprecated aliases for the old Metric-prefixed names -# --------------------------------------------------------------------------- - -def __getattr__(name: str) -> object: - _deprecated = { - "Metric": ("Evaluator", Evaluator), - "LlmMetric": ("LlmEvaluator", LlmEvaluator), - "CodeMetric": ("CodeEvaluator", CodeEvaluator), - "LocalMetric": ("LocalEvaluator", LocalEvaluator), - "SplunkAOMetric": ("SplunkAOEvaluator", SplunkAOEvaluator), - "BuiltInMetrics": ("BuiltInEvaluators", BuiltInEvaluators), - } - if name in _deprecated: - new_name, obj = _deprecated[name] - warnings.warn( - f"splunk_ao.evaluator.{name} is deprecated; use {new_name} instead.", - DeprecationWarning, - stacklevel=2, + if scorer_fn is None: + raise ValidationError("'scorer_fn' must be provided for local metrics.") + + self.scorer_fn = scorer_fn + self.scorable_types = scorable_types or [StepType.llm] + self.aggregatable_types = aggregatable_types or [StepType.trace] + self.scorer_type = None # Local metrics don't have a scorer_type + + def to_local_metric_config(self) -> LocalMetricConfig: + """ + Convert to LocalMetricConfig format. + + Returns + ------- + LocalMetricConfig: Local metric configuration for use with the logger. + + Examples + -------- + def my_scorer(trace): + return 0.5 + + metric = LocalEvaluator(name="test", scorer_fn=my_scorer) + config = metric.to_local_metric_config() + """ + return LocalMetricConfig( + name=self.name, + scorer_fn=self.scorer_fn, + scorable_types=self.scorable_types, + aggregatable_types=self.aggregatable_types, ) - return obj - raise AttributeError(f"module 'splunk_ao.evaluator' has no attribute {name!r}") + + def __repr__(self) -> str: + """Detailed string representation of the metric.""" + # Handle callables that don't have __name__ (partials, lambdas, callable instances) + fn_name = getattr(self.scorer_fn, "__name__", f"<{type(self.scorer_fn).__name__}>") + return f"LocalEvaluator(name='{self.name}', scorer_fn={fn_name})" diff --git a/src/splunk_ao/evaluators.py b/src/splunk_ao/evaluators.py index fcec5bed..cb4e08b5 100644 --- a/src/splunk_ao/evaluators.py +++ b/src/splunk_ao/evaluators.py @@ -1,53 +1,153 @@ -""" -Evaluators service layer — the renamed successor to Metrics (HYBIM-730). - -Provides ``Evaluators`` (service class) and the module-level convenience -functions ``create_custom_llm_evaluator``, ``get_evaluators``, and -``delete_evaluator``. - -The old ``metrics`` module and its symbols remain available but are deprecated. -""" -from __future__ import annotations - import datetime -import warnings +import logging -from splunk_ao.metrics import Metrics, create_custom_llm_metric, delete_metric, get_metrics +from galileo_core.schemas.logging.step import StepType +from splunk_ao.config import SplunkAOConfig +from splunk_ao.resources.api.data import ( + create_llm_scorer_version_scorers_scorer_id_version_llm_post, + create_scorers_post, + delete_scorer_scorers_scorer_id_delete, +) +from splunk_ao.resources.api.trace import query_metrics_projects_project_id_metrics_search_post +from splunk_ao.resources.models import ( + HTTPValidationError, + LogRecordsMetricsQueryRequest, + LogRecordsMetricsResponse, + ScorerTypes, +) from splunk_ao.resources.models.base_scorer_version_response import BaseScorerVersionResponse +from splunk_ao.resources.models.create_llm_scorer_version_request import CreateLLMScorerVersionRequest +from splunk_ao.resources.models.create_scorer_request import CreateScorerRequest from splunk_ao.resources.models.output_type_enum import OutputTypeEnum -from splunk_ao.resources.models.log_records_metrics_response import LogRecordsMetricsResponse -from galileo_core.schemas.logging.step import StepType +from splunk_ao.resources.models.scorer_defaults import ScorerDefaults +from splunk_ao.scorers import Scorers from splunk_ao.search import FilterType -__all__ = [ - "Evaluators", - "create_custom_llm_evaluator", - "delete_evaluator", - "get_evaluators", -] +_logger = logging.getLogger(__name__) -class Evaluators(Metrics): - """ - Low-level service class for managing evaluators. +class Evaluators: + config: SplunkAOConfig - ``Evaluators`` is the new name for the ``Metrics`` service class. - Inherits all methods from ``Metrics`` unchanged. + def __init__(self) -> None: + self.config = SplunkAOConfig.get() - Examples - -------- - from splunk_ao.evaluators import Evaluators + def delete_evaluator(self, name: str) -> None: + scorers_to_delete = Scorers().list(name=name) + if not scorers_to_delete: + raise ValueError(f"Scorer with name {name} not found.") - svc = Evaluators() - svc.delete_metric(name="old-evaluator") - """ + for scorer in scorers_to_delete: + response = delete_scorer_scorers_scorer_id_delete.sync(scorer_id=scorer.id, client=self.config.api_client) + + if isinstance(response, HTTPValidationError): + raise ValueError(response.detail) + if response is None: + raise ValueError("Failed to delete metric.") + + def create_custom_llm_evaluator( + self, + name: str, + user_prompt: str, + node_level: StepType = StepType.llm, + cot_enabled: bool = True, + model_name: str = "gpt-4.1-mini", + num_judges: int = 3, + description: str = "", + tags: list[str] | None = None, + output_type: OutputTypeEnum = OutputTypeEnum.BOOLEAN, + ground_truth: bool = False, + ) -> BaseScorerVersionResponse: + """ + Create a custom LLM metric. + + Parameters + ---------- + name: str + Name of the metric. + user_prompt: str + User prompt for the metric. + node_level: StepType + Node level for the metric. + cot_enabled: bool + Whether chain-of-thought is enabled. + model_name: str + Model name to use. + num_judges: int + Number of judges for the metric. + description: str + Description of the metric. + tags: List[str] + Tags associated with the metric. + output_type: OutputTypeEnum + Output type for the metric. + ground_truth: bool + Whether the scorer requires ground truth (``reference_output``) from the dataset. + When True, the judge LLM receives the row's ground-truth value in its prompt. + + Returns + ------- + BaseScorerVersionResponse + Response containing the created metric details. + """ + if tags is None: + tags = [] + create_scorer_request = CreateScorerRequest( + name=name, + scorer_type=ScorerTypes.LLM, + description=description, + tags=tags, + defaults=ScorerDefaults(model_name=model_name, num_judges=num_judges, cot_enabled=cot_enabled), + scoreable_node_types=[node_level], + output_type=output_type, + ground_truth=ground_truth, + ) + + scorer = create_scorers_post.sync(body=create_scorer_request, client=self.config.api_client) + version_req = CreateLLMScorerVersionRequest(user_prompt=user_prompt) + version_resp = create_llm_scorer_version_scorers_scorer_id_version_llm_post.sync( + scorer_id=scorer.id, body=version_req, client=self.config.api_client + ) + + _logger.info("Created custom LLM metric: %s", name) + + return version_resp + + def query( + self, + project_id: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + experiment_id: str | None = None, + log_stream_id: str | None = None, + filters: list[FilterType] | None = None, + group_by: str | None = None, + interval: int = 5, + ) -> LogRecordsMetricsResponse: + body = LogRecordsMetricsQueryRequest( + start_time=start_time, + end_time=end_time, + experiment_id=experiment_id, + log_stream_id=log_stream_id, + filters=filters or [], + group_by=group_by, + interval=interval, + ) + + response = query_metrics_projects_project_id_metrics_search_post.sync( + client=self.config.api_client, project_id=str(project_id), body=body + ) -# --------------------------------------------------------------------------- -# Module-level convenience functions -# --------------------------------------------------------------------------- + if isinstance(response, HTTPValidationError): + raise ValueError(response.detail) + if response is None: + raise ValueError("Failed to query for metrics.") + return response + +# Public functions def create_custom_llm_evaluator( name: str, user_prompt: str, @@ -61,50 +161,41 @@ def create_custom_llm_evaluator( ground_truth: bool = False, ) -> BaseScorerVersionResponse: """ - Create a custom LLM evaluator. - - This is the renamed equivalent of ``create_custom_llm_metric`` from - ``splunk_ao.metrics``. + Create a custom LLM metric. Parameters ---------- - name: - Name of the evaluator. - user_prompt: - Prompt template for the evaluator. - node_level: - Node level. Defaults to ``StepType.llm``. - cot_enabled: - Whether chain-of-thought reasoning is enabled. - model_name: - Model alias to use for judging. - num_judges: - Number of judge LLMs to use. - description: - Human-readable description. - tags: - Tags to associate with the evaluator. - output_type: - Output type (boolean, percentage, etc.). - ground_truth: - Whether the evaluator requires a ground-truth reference value. + name: str + Name of the metric. + user_prompt: str + User prompt for the metric. + node_level: StepType + Node level for the metric. + cot_enabled: bool + Whether chain-of-thought is enabled. + model_name: str + Model name to use. + num_judges: int + Number of judges for the metric. + description: str + Description of the metric. + tags: List[str] + Tags associated with the metric. + output_type: OutputTypeEnum + Output type for the metric. + ground_truth: bool + Whether the scorer requires ground truth (``reference_output``) from the dataset. + When True, the judge LLM receives the row's ground-truth value in its prompt. Returns ------- BaseScorerVersionResponse - The created evaluator version details. + Response containing the created metric details. """ - return create_custom_llm_metric( - name=name, - user_prompt=user_prompt, - node_level=node_level, - cot_enabled=cot_enabled, - model_name=model_name, - num_judges=num_judges, - description=description, - tags=tags, - output_type=output_type, - ground_truth=ground_truth, + if tags is None: + tags = [] + return Evaluators().create_custom_llm_evaluator( + name, user_prompt, node_level, cot_enabled, model_name, num_judges, description, tags, output_type, ground_truth ) @@ -113,46 +204,43 @@ def get_evaluators( start_time: datetime.datetime, end_time: datetime.datetime, experiment_id: str | None = None, - agent_stream_id: str | None = None, + log_stream_id: str | None = None, filters: list[FilterType] | None = None, group_by: str | None = None, interval: int = 5, ) -> LogRecordsMetricsResponse: - """ - Query evaluator results for a project. - - This is the renamed equivalent of ``get_metrics`` from ``splunk_ao.metrics``. + """Queries for metrics in a project. Parameters ---------- - project_id: - Project UUID. - start_time: - Start of the query window. - end_time: - End of the query window. - experiment_id: - Filter by experiment ID (optional). - agent_stream_id: - Filter by agent stream ID (optional). - filters: - Additional query filters. - group_by: - Field to group results by. - interval: - Time interval in seconds. + project_id + The unique identifier of the project. + start_time + The start of the time range for the query. + end_time + The end of the time range for the query. + experiment_id + Filter records by a specific experiment ID. + log_stream_id + Filter records by a specific run ID. + filters + A list of filters to apply to the query. + group_by + The field to group the results by. + interval + The time interval for the query in seconds. Returns ------- LogRecordsMetricsResponse - Evaluator query results. + A LogRecordsMetricsResponse object containing the query results, or None if the query fails. """ - return get_metrics( + return Evaluators().query( project_id=project_id, start_time=start_time, end_time=end_time, experiment_id=experiment_id, - log_stream_id=agent_stream_id, + log_stream_id=log_stream_id, filters=filters, group_by=group_by, interval=interval, @@ -161,34 +249,11 @@ def get_evaluators( def delete_evaluator(name: str) -> None: """ - Delete an evaluator by name. - - This is the renamed equivalent of ``delete_metric`` from ``splunk_ao.metrics``. + Deletes a metric by its name. Parameters ---------- - name: - The evaluator name to delete. + name + The name of the metric to delete. """ - return delete_metric(name=name) - - -# --------------------------------------------------------------------------- -# Deprecated aliases for old ``metrics`` module function names -# --------------------------------------------------------------------------- - -def __getattr__(name: str) -> object: - _deprecated = { - "create_custom_llm_metric": ("create_custom_llm_evaluator", create_custom_llm_evaluator), - "get_metrics": ("get_evaluators", get_evaluators), - "delete_metric": ("delete_evaluator", delete_evaluator), - } - if name in _deprecated: - new_name, obj = _deprecated[name] - warnings.warn( - f"splunk_ao.evaluators.{name} is deprecated; use {new_name} instead.", - DeprecationWarning, - stacklevel=2, - ) - return obj - raise AttributeError(f"module 'splunk_ao.evaluators' has no attribute {name!r}") + Evaluators().delete_evaluator(name) diff --git a/src/splunk_ao/export.py b/src/splunk_ao/export.py index 8c2319dd..5d4ee819 100644 --- a/src/splunk_ao/export.py +++ b/src/splunk_ao/export.py @@ -6,7 +6,7 @@ from typing import Any from splunk_ao.config import SplunkAOConfig -from splunk_ao.log_streams import LogStreams +from splunk_ao.agent_streams import AgentStreams from splunk_ao.resources.api.trace.export_records_projects_project_id_export_records_post import ( stream_detailed as export_records_stream, ) @@ -121,7 +121,7 @@ def export_records( if log_stream_id is None and experiment_id is None: # Use _list_all to paginate across all pages so we pick the globally oldest # stream, not just the oldest in the first page (default page size is 100). - log_streams = LogStreams()._list_all(project_id=project_id) + log_streams = AgentStreams()._list_all(project_id=project_id) if log_streams: sorted_log_streams = sorted(log_streams, key=lambda ls: (ls.created_at, ls.id)) log_stream_id = sorted_log_streams[0].id diff --git a/src/splunk_ao/log_stream.py b/src/splunk_ao/log_stream.py deleted file mode 100644 index 81afefb4..00000000 --- a/src/splunk_ao/log_stream.py +++ /dev/null @@ -1,971 +0,0 @@ -from __future__ import annotations - -import builtins -import logging -from collections.abc import Iterator -from datetime import datetime -from typing import TYPE_CHECKING, Any - -from splunk_ao.config import SplunkAOConfig -from splunk_ao.decorator import splunk_ao_context -from splunk_ao.export import ExportClient -from splunk_ao.log_streams import LogStreams -from splunk_ao.resources.api.trace import ( - sessions_available_columns_projects_project_id_sessions_available_columns_post, - spans_available_columns_projects_project_id_spans_available_columns_post, - traces_available_columns_projects_project_id_traces_available_columns_post, -) -from splunk_ao.resources.models import LLMExportFormat, LogRecordsSortClause, RootType -from splunk_ao.resources.models.http_validation_error import HTTPValidationError -from splunk_ao.resources.models.log_records_available_columns_request import LogRecordsAvailableColumnsRequest -from splunk_ao.resources.models.log_records_available_columns_response import LogRecordsAvailableColumnsResponse -from splunk_ao.resources.types import Unset -from splunk_ao.schema.filters import FilterType -from splunk_ao.schema.metrics import LocalMetricConfig, Metric, SplunkAOMetrics -from splunk_ao.search import RecordType, Search -from splunk_ao.shared.base import StateManagementMixin, SyncState -from splunk_ao.shared.exceptions import ValidationError -from splunk_ao.shared.project_resolver import _resolve_project -from splunk_ao.shared.query_result import QueryResult - -if TYPE_CHECKING: - from splunk_ao.shared.column import ColumnCollection - -logger = logging.getLogger(__name__) - -# Mapping from RecordType (plural) to RootType (singular) -RECORD_TYPE_TO_ROOT_TYPE = { - RecordType.SPAN: RootType.SPAN, - RecordType.TRACE: RootType.TRACE, - RecordType.SESSION: RootType.SESSION, -} - - -class LogStream(StateManagementMixin): - """ - Object-centric interface for Galileo log streams. - - This class provides an intuitive way to work with Galileo log streams, - offering methods for managing log streams and their associated metrics. - - Attributes - ---------- - created_at (datetime.datetime): When the log stream was created. - created_by (str): The user who created the log stream. - id (str): The unique log stream identifier. - name (str): The log stream name. - project_id (str): The ID of the project this log stream belongs to. - project_name (str | None): The name of the project. May be None if the log stream - was retrieved using project_id only, as the API doesn't return this field. - updated_at (datetime.datetime): When the log stream was last updated. - additional_properties (dict): Additional properties of the log stream. - - Examples - -------- - # Create a new log stream and persist it - log_stream = LogStream(name="Production Logs", project_name="My AI Project").create() - - # Get an existing log stream - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - - # LogStreams can also be created through Project instances - from splunk_ao.project import Project - - project = Project.get(name="My AI Project") - log_stream = project.create_log_stream(name="Production Logs") - - # Enable metrics on the log stream - from splunk_ao.schema.metrics import SplunkAOMetrics - local_metrics = log_stream.enable_metrics([ - SplunkAOMetrics.correctness, - SplunkAOMetrics.completeness, - "context_relevance" - ]) - - # Refresh log stream state from API - log_stream.refresh() - """ - - created_at: datetime | None - created_by: str | None - id: str | None - name: str - project_id: str | None - project_name: str | None # May not be available when retrieved from API - updated_at: datetime | None - additional_properties: dict[str, Any] # TODO: We need to validate if we will keep this one. - - def __str__(self) -> str: - """String representation of the log stream.""" - return f"LogStream(name='{self.name}', id='{self.id}', project_id='{self.project_id}')" - - def __repr__(self) -> str: - """Detailed string representation of the log stream.""" - return f"LogStream(name='{self.name}', id='{self.id}', project_id='{self.project_id}', created_at='{self.created_at}')" - - def __init__(self, name: str, *, project_id: str | None = None, project_name: str | None = None) -> None: - """ - Initialize a LogStream instance locally. - - Creates a local log stream object that exists only in memory until .create() - is called to persist it to the API. - - Args: - name (str): The name of the log stream to create. - project_id (Optional[str]): The project ID. If neither project_id nor project_name is provided, - falls back to SPLUNK_AO_PROJECT_ID or SPLUNK_AO_PROJECT environment variables. - project_name (Optional[str]): The project name. If neither project_id nor project_name is provided, - falls back to SPLUNK_AO_PROJECT environment variable. - - Raises - ------ - ValidationError: If name is not provided. - ValueError: If project cannot be resolved at create() time (no explicit param and no env fallback). - - Examples - -------- - # Create by project ID - log_stream = LogStream(name="Production Logs", project_id="project-123") - - # Create by project name - log_stream = LogStream(name="Production Logs", project_name="My AI Project") - - # Create using SPLUNK_AO_PROJECT environment variable - log_stream = LogStream(name="Production Logs") - """ - super().__init__() - - if not name: - raise ValidationError("'name' must be provided to create a log stream.") - - # Initialize attributes locally - self.name = name - self.project_id = project_id - self.project_name = project_name # May be None; not returned by API - self.id = None - self.created_at = None - self.created_by = None - self.updated_at = None - self.additional_properties = {} - - # Set initial state - self._set_state(SyncState.LOCAL_ONLY) - - def create(self) -> LogStream: - """ - Persist this log stream to the API. - - Returns - ------- - LogStream: This log stream instance with updated attributes from the API. - - Raises - ------ - NotFoundError: If the project cannot be found (no explicit param and no env fallback). - Exception: If the API call fails. - - Examples - -------- - log_stream = LogStream(name="Production Logs", project_name="My AI Project").create() - assert log_stream.is_synced() - """ - if not self.name: - raise ValueError("Log stream name is not set. Cannot create log stream without a name.") - - # Note: project_id and project_name can both be None here — resolution happens below - # via _resolve_project, which reads SPLUNK_AO_PROJECT_ID / SPLUNK_AO_PROJECT env vars. - - try: - logger.info(f"LogStream.create: name='{self.name}' project_id='{self.project_id}' - started") - - # _resolve_project raises NotFoundError when no project can be found, matching - # the contract of LogStream.get() and LogStream.list(). - project_obj = _resolve_project(self.project_id, self.project_name) - - # Update project info from resolved project - self.project_id = project_obj.id - if self.project_name is None: - self.project_name = project_obj.name - - log_streams_service = LogStreams() - created_log_stream = log_streams_service.create( - name=self.name, - project_id=self.project_id, - project_name=None, # Use resolved project_id - ) - - # Update attributes from response - self.created_at = created_log_stream.created_at - self.created_by = created_log_stream.created_by - self.id = created_log_stream.id - self.name = created_log_stream.name - self.project_id = created_log_stream.project_id - self.updated_at = created_log_stream.updated_at - self.additional_properties = created_log_stream.additional_properties - # Note: project_name is preserved if it was set, but API doesn't return it - - # Set state to synced - self._set_state(SyncState.SYNCED) - logger.info(f"LogStream.create: id='{self.id}' - completed") - return self - except Exception as e: - self._set_state(SyncState.FAILED_SYNC, error=e) - logger.error(f"LogStream.create: name='{self.name}' - failed: {e}") - raise - - @classmethod - def _create_empty(cls) -> LogStream: - """Internal constructor bypassing __init__ for API hydration.""" - instance = cls.__new__(cls) - super(LogStream, instance).__init__() - return instance - - @classmethod - def _from_api_response(cls, retrieved_log_stream: Any) -> LogStream: - """ - Factory method to create a LogStream instance from an API response. - - Args: - retrieved_log_stream: The log stream data retrieved from the API. - - Returns - ------- - LogStream: A new LogStream instance populated with the API data. - """ - instance = cls._create_empty() - instance.created_at = retrieved_log_stream.created_at - instance.created_by = retrieved_log_stream.created_by - instance.id = retrieved_log_stream.id - instance.name = retrieved_log_stream.name - instance.project_id = retrieved_log_stream.project_id - instance.updated_at = retrieved_log_stream.updated_at - instance.additional_properties = retrieved_log_stream.additional_properties - instance.project_name = None # API doesn't return project_name - # Set state to synced since we just retrieved from API - instance._set_state(SyncState.SYNCED) - return instance - - @classmethod - def get(cls, *, name: str, project_id: str | None = None, project_name: str | None = None) -> LogStream | None: - """ - Get an existing log stream by name. - - Args: - name (str): The log stream name. - project_id (Optional[str]): The project ID. If neither project_id nor project_name is provided, - falls back to SPLUNK_AO_PROJECT_ID or SPLUNK_AO_PROJECT environment variables. - project_name (Optional[str]): The project name. If neither project_id nor project_name is provided, - falls back to SPLUNK_AO_PROJECT environment variable. - - Returns - ------- - Optional[LogStream]: The log stream if found, None otherwise. - - Raises - ------ - NotFoundError: If the project cannot be found (no explicit param and no env fallback). - - Examples - -------- - # Get by project name - log_stream = LogStream.get( - name="Production Logs", - project_name="My AI Project" - ) - - # Get by project ID - log_stream = LogStream.get( - name="Production Logs", - project_id="project-123" - ) - - # Get using SPLUNK_AO_PROJECT environment variable - log_stream = LogStream.get(name="Production Logs") - """ - project_obj = _resolve_project(project_id, project_name) - - log_streams_service = LogStreams() - retrieved_log_stream = log_streams_service.get(name=name, project_id=project_obj.id) - if retrieved_log_stream is None: - return None - - instance = cls._from_api_response(retrieved_log_stream) - # Set project_name from resolved project - instance.project_name = project_obj.name - return instance - - @classmethod - def list( - cls, - *, - project_id: str | None = None, - project_name: str | None = None, - limit: Unset | int = 100, - starting_token: Unset | int = 0, - ) -> list[LogStream]: - """ - List log streams for a project. - - Returns a single page of results. Use `starting_token` (from - `next_starting_token` on a prior response) to fetch subsequent pages. - - Args: - project_id (Optional[str]): The project ID. If neither project_id nor project_name is provided, - falls back to SPLUNK_AO_PROJECT_ID or SPLUNK_AO_PROJECT environment variables. - project_name (Optional[str]): The project name. If neither project_id nor project_name is provided, - falls back to SPLUNK_AO_PROJECT environment variable. - limit (Union[Unset, int]): Maximum number of log streams to return per page. Defaults to 100. - starting_token (Union[Unset, int]): Pagination token to start from. Defaults to 0 (first page). - - Returns - ------- - List[LogStream]: A page of log streams for the project. - - Raises - ------ - NotFoundError: If the project cannot be found (no explicit param and no env fallback). - - Examples - -------- - # List by project name - log_streams = LogStream.list(project_name="My AI Project") - - # List by project ID - log_streams = LogStream.list(project_id="project-123") - - # List using SPLUNK_AO_PROJECT environment variable - log_streams = LogStream.list() - - # Cap the number of returned log streams - log_streams = LogStream.list(project_name="My AI Project", limit=3) - - # Fetch the next page - page_2 = LogStream.list(project_name="My AI Project", starting_token=100) - """ - project_obj = _resolve_project(project_id, project_name) - - log_streams_service = LogStreams() - retrieved_log_streams = log_streams_service.list( - project_id=project_obj.id, limit=limit, starting_token=starting_token - ) - - instances = [cls._from_api_response(retrieved_log_stream) for retrieved_log_stream in retrieved_log_streams] - # Set project_name from resolved project for all instances - for instance in instances: - instance.project_name = project_obj.name - return instances - - def refresh(self) -> None: - """ - Refresh this log stream's state from the API. - - Updates all attributes with the latest values from the remote API - and sets the state to SYNCED. - - Raises - ------ - ValueError: If the log stream ID or project_id is not set. - Exception: If the API call fails or the log stream no longer exists. - - Examples - -------- - log_stream.refresh() - assert log_stream.is_synced() - """ - if self.id is None: - raise ValueError("Log stream ID is not set. Cannot refresh a local-only log stream.") - - if self.project_id is None: - raise ValueError("Project ID is not set. Cannot refresh log stream without project_id.") - - try: - logger.debug(f"LogStream.refresh: id='{self.id}' - started") - log_streams_service = LogStreams() - retrieved_log_stream = log_streams_service.get(id=self.id, project_id=self.project_id) - - if retrieved_log_stream is None: - raise ValueError(f"Log stream with id '{self.id}' no longer exists") - - # Update all attributes from response - self.created_at = retrieved_log_stream.created_at - self.created_by = retrieved_log_stream.created_by - self.id = retrieved_log_stream.id - self.name = retrieved_log_stream.name - self.project_id = retrieved_log_stream.project_id - self.updated_at = retrieved_log_stream.updated_at - self.additional_properties = retrieved_log_stream.additional_properties - # Note: project_name is preserved from before refresh since API doesn't return it - - # Set state to synced - self._set_state(SyncState.SYNCED) - logger.debug(f"LogStream.refresh: id='{self.id}' - completed") - except Exception as e: - self._set_state(SyncState.FAILED_SYNC, error=e) - logger.error(f"LogStream.refresh: id='{self.id}' - failed: {e}") - raise - - def get_metrics(self) -> builtins.list[str]: - """ - Get the list of metrics currently enabled on this log stream. - - Returns - ------- - list[str]: List of metric names currently enabled. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - log_stream = LogStream.get(name="Production Logs", project_name="My Project") - current_metrics = log_stream.get_metrics() - print(f"Currently enabled: {current_metrics}") - """ - logger.info(f"LogStream.get_metrics: id='{self.id}' - started") - config = SplunkAOConfig.get() - - settings = get_settings_projects_project_id_runs_run_id_scorer_settings_get.sync( - project_id=self.project_id, run_id=self.id, client=config.api_client - ) - - if settings is None or not hasattr(settings, "scorers"): - logger.info(f"LogStream.get_metrics: id='{self.id}' - no metrics enabled") - return [] - - # Extract metric names from scorer configs - metric_names = [scorer.name for scorer in settings.scorers] - logger.info(f"LogStream.get_metrics: id='{self.id}' found {len(metric_names)} metrics - completed") - return metric_names - - def set_metrics( - self, metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str] - ) -> builtins.list[LocalMetricConfig]: - """ - Set (replace) the metrics on this log stream. - - This replaces any existing metrics with the new list. Alias for enable_metrics - with clearer naming intent. - - Args: - metrics: List of metrics to set. Supports: - - SplunkAOMetrics enum values (e.g., SplunkAOMetrics.correctness) - - Metric objects (including from Metric.get(id="...")) - - LocalMetricConfig objects for custom scoring functions - - String names of built-in metrics - - Returns - ------- - List[LocalMetricConfig]: Local metric configurations that must be - computed client-side. - - Raises - ------ - ValueError: If any specified metrics are unknown. - - Examples - -------- - from splunk_ao import Metric, LogStream - - log_stream = LogStream.get(name="Production Logs", project_name="My Project") - - # Set metrics (replaces existing) - log_stream.set_metrics([ - Metric.metrics.correctness, - Metric.metrics.completeness, - Metric.get(id="metric-from-console-uuid"), # From console - ]) - """ - try: - logger.info(f"LogStream.enable_metrics: id='{self.id}' metrics={[str(m) for m in metrics]} - started") - log_streams_service = LogStreams() - log_stream = log_streams_service.get(name=self.name, project_id=self.project_id) - if log_stream is None: - raise ValueError(f"Log stream '{self.name}' not found") - result = log_stream.enable_metrics(metrics) - # Set state to synced after successful operation - self._set_state(SyncState.SYNCED) - logger.info(f"LogStream.enable_metrics: id='{self.id}' - completed") - return result - except Exception as e: - self._set_state(SyncState.FAILED_SYNC, error=e) - logger.error(f"LogStream.enable_metrics: id='{self.id}' - failed: {e}") - raise - - def query( - self, - record_type: RecordType, - filters: builtins.list[FilterType] | None = None, - sort: LogRecordsSortClause | None = None, - limit: int = 100, - starting_token: int = 0, - ) -> QueryResult: - """ - Query records in this log stream. - - This method provides a convenient way to search spans, traces, or sessions - within the current log stream without needing to specify project_id or log_stream_id. - - Args: - record_type: The type of records to query (SPAN, TRACE, or SESSION). - filters: A list of filters to apply to the query. - sort: A sort clause to order the query results. - limit: The maximum number of records to return. - starting_token: The token for the next page of results. - - Returns - ------- - QueryResult: A list-like object containing the query results with pagination support. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - from splunk_ao.search import RecordType - - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - - # Query with column-based filters and sort - results = log_stream.query( - record_type=RecordType.SPAN, - filters=[ - log_stream.span_columns["input"].contains("largest"), - log_stream.span_columns["metrics/completeness_gpt"].greater_than(0.8), - log_stream.span_columns["created_at"].after("2024-01-01") - ], - sort=log_stream.span_columns["created_at"].descending(), - limit=50 - ) - - # Access results like a list - for record in results: - print(record["id"], record["input"]) - - # Get specific record - first_record = results[0] - - # Pagination - if results.has_next_page: - next_results = results.next_page() - """ - if self.id is None: - raise ValueError("Log stream ID is not set. Cannot query a local-only log stream.") - if self.project_id is None: - raise ValueError("Project ID is not set. Cannot query log stream without project_id.") - - logger.debug(f"LogStream.query: id='{self.id}' record_type='{record_type.value}' limit={limit} - started") - - # Capture project_id and log_stream_id for use in pagination function - project_id = self.project_id - log_stream_id = self.id - - search_service = Search() - response = search_service.query( - project_id=project_id, - record_type=record_type, - log_stream_id=log_stream_id, - filters=filters, - sort=sort, - limit=limit, - starting_token=starting_token, - ) - - # Create a query function that returns raw response for pagination - def query_fn( - record_type: RecordType, - filters: builtins.list[FilterType] | None, - sort: LogRecordsSortClause | None, - limit: int, - starting_token: int, - ) -> Any: - return Search().query( - project_id=project_id, - record_type=record_type, - log_stream_id=log_stream_id, - filters=filters, - sort=sort, - limit=limit, - starting_token=starting_token, - ) - - # Wrap the response in QueryResult for easy access and pagination - return QueryResult(response=response, query_fn=query_fn, record_type=record_type, filters=filters, sort=sort) - - def get_spans( - self, - filters: builtins.list[FilterType] | None = None, - sort: LogRecordsSortClause | None = None, - limit: int = 100, - starting_token: int = 0, - ) -> QueryResult: - """ - Query spans in this log stream. - - This is a convenience method that queries for spans specifically. - - Args: - filters: A list of filters to apply to the query. - sort: A sort clause to order the query results. - limit: The maximum number of records to return. - starting_token: The token for the next page of results. - - Returns - ------- - QueryResult: A list-like object containing the span query results with pagination support. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - - # Get spans with filters and sorting - spans = log_stream.get_spans( - filters=[ - log_stream.span_columns["input"].contains("world"), - log_stream.span_columns["metrics/num_input_tokens"].greater_than(10) - ], - sort=log_stream.span_columns["created_at"].descending(), - limit=50 - ) - - # Iterate over results - for span in spans: - print(span["id"], span["input"]) - - # Pagination - if spans.has_next_page: - more_spans = spans.next_page() - """ - logger.debug(f"LogStream.get_spans: id='{self.id}' limit={limit} - started") - return self.query( - record_type=RecordType.SPAN, filters=filters, sort=sort, limit=limit, starting_token=starting_token - ) - - def get_traces( - self, - filters: builtins.list[FilterType] | None = None, - sort: LogRecordsSortClause | None = None, - limit: int = 100, - starting_token: int = 0, - ) -> QueryResult: - """ - Query traces in this log stream. - - This is a convenience method that queries for traces specifically. - - Args: - filters: A list of filters to apply to the query. - sort: A sort clause to order the query results. - limit: The maximum number of records to return. - starting_token: The token for the next page of results. - - Returns - ------- - QueryResult: A list-like object containing the trace query results with pagination support. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - - # Get traces with filters - traces = log_stream.get_traces( - filters=[ - log_stream.trace_columns["input"].contains("largest"), - log_stream.trace_columns["created_at"].after("2024-01-01") - ], - sort=log_stream.trace_columns["created_at"].descending(), - limit=50 - ) - - # Access like a list - for trace in traces: - print(trace["id"], trace["input"]) - """ - logger.debug(f"LogStream.get_traces: id='{self.id}' limit={limit} - started") - return self.query( - record_type=RecordType.TRACE, filters=filters, sort=sort, limit=limit, starting_token=starting_token - ) - - def get_sessions( - self, - filters: builtins.list[FilterType] | None = None, - sort: LogRecordsSortClause | None = None, - limit: int = 100, - starting_token: int = 0, - ) -> QueryResult: - """ - Query sessions in this log stream. - - This is a convenience method that queries for sessions specifically. - - Args: - filters: A list of filters to apply to the query. - sort: A sort clause to order the query results. - limit: The maximum number of records to return. - starting_token: The token for the next page of results. - - Returns - ------- - QueryResult: A list-like object containing the session query results with pagination support. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - - # Get sessions with filters - sessions = log_stream.get_sessions( - filters=[ - log_stream.session_columns["model"].equals("gpt-4o-mini"), - log_stream.session_columns["metrics/num_traces"].greater_than(5) - ], - sort=log_stream.session_columns["created_at"].descending(), - limit=50 - ) - - # Work with results - for session in sessions: - print(session["id"], session["model"]) - """ - logger.debug(f"LogStream.get_sessions: id='{self.id}' limit={limit} - started") - return self.query( - record_type=RecordType.SESSION, filters=filters, sort=sort, limit=limit, starting_token=starting_token - ) - - def export_records( - self, - record_type: RecordType = RecordType.TRACE, - filters: builtins.list[FilterType] | None = None, - sort: LogRecordsSortClause = LogRecordsSortClause(column_id="created_at", ascending=False), - export_format: LLMExportFormat = LLMExportFormat.JSONL, - column_ids: builtins.list[str] | None = None, - redact: bool = True, - ) -> Iterator[dict[str, Any]]: - """ - Export records from this log stream. - - This method provides a convenient way to export records without needing - to specify project_id or log_stream_id. - - Args: - record_type: The type of records to export (SPAN, TRACE, or SESSION). - filters: A list of filters to apply to the export. - sort: A sort clause to order the exported records. - export_format: The desired format for the exported data. - column_ids: A list of column IDs to include in the export. - redact: Redact sensitive data from the response. - - Returns - ------- - Iterator[dict[str, Any]]: An iterator that yields each record as a dictionary. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - from splunk_ao.search import RecordType - - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - - # Export records with filters - for record in log_stream.export_records( - record_type=RecordType.SPAN, - filters=[ - log_stream.span_columns["model"].one_of(["gpt-4", "gpt-3.5-turbo", "gpt-4o-mini"]), - log_stream.span_columns["metrics/num_input_tokens"].greater_than(1) - ], - sort=log_stream.span_columns["created_at"].descending() - ): - print(record) - """ - if self.id is None: - raise ValueError("Log stream ID is not set. Cannot export from a local-only log stream.") - if self.project_id is None: - raise ValueError("Project ID is not set. Cannot export log stream without project_id.") - - # Convert RecordType to RootType for the export client - root_type = RECORD_TYPE_TO_ROOT_TYPE[record_type] - - logger.info( - f"LogStream.export_records: id='{self.id}' record_type='{record_type.value}' " - f"export_format='{export_format.value}' - started" - ) - export_client = ExportClient() - return export_client.records( - project_id=self.project_id, - root_type=root_type, - filters=filters, - sort=sort, - export_format=export_format, - log_stream_id=self.id, - column_ids=column_ids, - redact=redact, - ) - - def context(self) -> Any: - """ - Get a galileo context manager for this log stream. - - This is a convenient method that returns a pre-configured splunk_ao_context - for this log stream, eliminating the need to specify project and log stream names. - - Returns - ------- - A context manager for Galileo logging configured with this log stream. - - Examples - -------- - log_stream = LogStream.get( - name="Production Logs", - project_name="My AI Project" - ) - - with log_stream.context(): - # Your logging code here - response = openai_client.chat.completions.create(...) - """ - return splunk_ao_context(project=self.project.name if self.project else None, log_stream=self.name) - - def _get_columns(self, api_func: Any, error_msg: str) -> LogRecordsAvailableColumnsResponse: - """Helper method to retrieve available columns from the API.""" - if self.id is None: - raise ValueError("Log stream ID is not set. Cannot get columns from a local-only log stream.") - if self.project_id is None: - raise ValueError("Project ID is not set. Cannot get columns without project_id.") - - config = SplunkAOConfig.get() - body = LogRecordsAvailableColumnsRequest(log_stream_id=self.id) - response = api_func.sync(project_id=self.project_id, client=config.api_client, body=body) - if isinstance(response, HTTPValidationError): - raise response - if not response: - raise ValueError(error_msg) - return response - - @property - def project(self) -> Project | None: - """Get the project this log stream belongs to.""" - return Project.get(id=self.project_id) - - @property - def span_columns(self) -> ColumnCollection: - """ - Get available columns for spans in this log stream. - - Returns - ------- - ColumnCollection: A collection of columns available for spans, accessible by column ID. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - columns = log_stream.span_columns - - # Access a specific column - input_column = columns["input"] - - # Filter using columns - spans = log_stream.get_spans( - filters=[columns["input"].contains("world")], - sort=columns["created_at"].descending() - ) - """ - response = self._get_columns( - spans_available_columns_projects_project_id_spans_available_columns_post, "Unable to retrieve span columns" - ) - columns = [Column(col) for col in response.columns] - return ColumnCollection(columns) - - @property - def session_columns(self) -> ColumnCollection: - """ - Get available columns for sessions in this log stream. - - Returns - ------- - ColumnCollection: A collection of columns available for sessions, accessible by column ID. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - columns = log_stream.session_columns - - # Access a specific column - model_column = columns["model"] - - # Filter using columns - sessions = log_stream.get_sessions( - filters=[columns["model"].equals("gpt-4o-mini")], - sort=columns["created_at"].descending() - ) - """ - response = self._get_columns( - sessions_available_columns_projects_project_id_sessions_available_columns_post, - "Unable to retrieve session columns", - ) - columns = [Column(col) for col in response.columns] - return ColumnCollection(columns) - - @property - def trace_columns(self) -> ColumnCollection: - """ - Get available columns for traces in this log stream. - - Returns - ------- - ColumnCollection: A collection of columns available for traces, accessible by column ID. - - Raises - ------ - ValueError: If the log stream lacks required id or project_id attributes. - - Examples - -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") - columns = log_stream.trace_columns - - # Access a specific column - input_column = columns["input"] - - # Filter using columns - traces = log_stream.get_traces( - filters=[columns["input"].contains("largest")], - sort=columns["created_at"].descending() - ) - """ - response = self._get_columns( - traces_available_columns_projects_project_id_traces_available_columns_post, - "Unable to retrieve trace columns", - ) - columns = [Column(col) for col in response.columns] - return ColumnCollection(columns) - - -# Import at end to avoid circular import (project.py imports LogStream) -from splunk_ao.project import Project # noqa: E402 -from splunk_ao.resources.api.run_scorer_settings import ( # noqa: E402 - get_settings_projects_project_id_runs_run_id_scorer_settings_get, -) -from splunk_ao.shared.column import Column, ColumnCollection # noqa: E402 diff --git a/src/splunk_ao/log_streams.py b/src/splunk_ao/log_streams.py deleted file mode 100644 index 662a646c..00000000 --- a/src/splunk_ao/log_streams.py +++ /dev/null @@ -1,790 +0,0 @@ -import builtins -from typing import overload - -from splunk_ao.config import SplunkAOConfig -from splunk_ao.projects import Projects -from splunk_ao.resources.api.log_stream import ( - create_log_stream_projects_project_id_log_streams_post, - get_log_stream_projects_project_id_log_streams_log_stream_id_get, - list_log_streams_paginated_projects_project_id_log_streams_paginated_get, -) -from splunk_ao.resources.models.http_validation_error import HTTPValidationError -from splunk_ao.resources.models.log_stream_create_request import LogStreamCreateRequest -from splunk_ao.resources.models.log_stream_response import LogStreamResponse -from splunk_ao.resources.types import Unset -from splunk_ao.schema.metrics import LocalMetricConfig, Metric, SplunkAOMetrics -from splunk_ao.utils.env_helpers import _get_log_stream_from_env, _get_project_from_env -from splunk_ao.utils.log_config import get_logger -from splunk_ao.utils.metrics import create_metric_configs - -logger = get_logger(__name__) - - -class LogStream(LogStreamResponse): - """ - Log streams are used to organize logs within a project on the Galileo platform. - They provide a way to categorize and group related logs, making it easier to - analyze and monitor specific parts of your application or different environments - (e.g., production, staging, development). - - Attributes - ---------- - created_at : datetime.datetime - The timestamp when the log stream was created. - created_by : str - The identifier of the user who created the log stream. - id : str - The unique identifier of the log stream. - name : str - The name of the log stream. - project_id : str - The ID of the project this log stream belongs to. - updated_at : datetime.datetime - The timestamp when the log stream was last updated. - additional_properties : dict - Additional properties associated with the log stream. - - Examples - -------- - ```python - # Create a new log stream in a project - from splunk_ao.log_streams import create_log_stream - - # Create by project ID - log_stream = create_log_stream(name="Production Logs", project_id="project-123") - - # Create by project name - log_stream = create_log_stream(name="Production Logs", project_name="My AI Project") - - # Get a log stream by name - from splunk_ao.log_streams import get_log_stream - log_stream = get_log_stream(name="Production Logs", project_name="My AI Project") - - # List all log streams in a project - from splunk_ao.log_streams import list_log_streams - log_streams = list_log_streams(project_name="My AI Project") - for stream in log_streams: - logger.info(f"Log Stream: {stream.name} (ID: {stream.id})") - - # Use a log stream with the context manager - from splunk_ao.openai import openai - from splunk_ao import splunk_ao_context - - with splunk_ao_context(project="My AI Project", log_stream="Production Logs"): - response = openai.chat.completions.create( - model="gpt-4o", - messages=[{"role": "user", "content": "Hello, world!"}] - ) - - # Enable metrics on a log stream - RECOMMENDED APPROACH - from splunk_ao.log_streams import enable_metrics - from splunk_ao.schema.metrics import SplunkAOMetrics - - # Set environment variables first - # export SPLUNK_AO_LOG_STREAM="Production Logs" - # export SPLUNK_AO_PROJECT="My AI Project" - - # Clean and simple - just pass the metrics! - local_metrics = enable_metrics([ - SplunkAOMetrics.correctness, - SplunkAOMetrics.completeness, - "context_relevance" - ]) - - # Alternative: Use explicit parameters - local_metrics = enable_metrics( - log_stream_name="Production Logs", - project_name="My AI Project", - metrics=["correctness", "completeness"] - ) - ``` - """ - - def __init__(self, log_stream: None | LogStreamResponse = None): - """ - Initialize a LogStream instance. - - Parameters - ---------- - log_stream : Union[None, LogStreamResponse], optional - The log stream data to initialize from. If None, creates an empty log stream instance. - Defaults to None. - """ - if log_stream is not None: - super().__init__( - created_at=log_stream.created_at, - id=log_stream.id, - name=log_stream.name, - project_id=log_stream.project_id, - updated_at=log_stream.updated_at, - created_by=log_stream.created_by, - ) - self.additional_properties = log_stream.additional_properties.copy() - return - - def enable_metrics( - self, metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str] - ) -> builtins.list[LocalMetricConfig]: - """ - Enable metrics directly on this log stream instance. - - This is the most intuitive and clean way to enable metrics when you already have a - LogStream object. The method leverages the log stream's existing project_id and id - attributes, eliminating the need for redundant parameter specification and reducing - the potential for errors. - - This approach is ideal for object-oriented workflows where you're working with LogStream - instances directly, and it provides the clearest semantic meaning: "enable these metrics - on this specific log stream." - - Parameters - ---------- - metrics : builtins.list[Union[SplunkAOMetrics, Metric, LocalMetricConfig, str]] - List of metrics to enable on this log stream. Supports multiple input formats: - - - **SplunkAOMetrics enum values**: Built-in metrics like `SplunkAOMetrics.correctness` - - **Metric objects**: Custom metrics with optional version specifications - - **LocalMetricConfig objects**: Client-side metrics with custom scoring functions - - **String names**: Built-in metric names like "correctness" or "toxicity" - - Returns - ------- - builtins.list[LocalMetricConfig] - List of local metric configurations that must be computed client-side. - Server-side metrics are automatically registered with Galileo and don't - need to be returned since users don't interact with them. - - Raises - ------ - ValueError - - If this LogStream instance lacks required `id` or `project_id` attributes - - If any specified metrics are unknown or unavailable - - If there are issues with metric configuration or registration - GalileoHTTPException - If there are network or API errors when communicating with Galileo services - - Examples - -------- - Basic usage with built-in metrics: - - ```python - from splunk_ao.log_streams import LogStreams - from splunk_ao.schema.metrics import SplunkAOMetrics - - # Get a log stream first - log_streams = LogStreams() - log_stream = log_streams.get(name="Production Logs", project_name="My AI Project") - - # Enable metrics directly - clean and intuitive! - local_metrics = log_stream.enable_metrics([ - SplunkAOMetrics.correctness, - SplunkAOMetrics.completeness, - "context_relevance", - "toxicity" - ]) - - logger.info(f"Server-side metrics enabled automatically") - logger.info(f"Need to process {len(local_metrics)} local metrics") - ``` - - Advanced usage with custom metrics: - - ```python - from splunk_ao.schema.metrics import Metric, LocalMetricConfig - - def custom_scorer(trace_or_span): - return 0.75 # Your scoring logic - - local_metrics = log_stream.enable_metrics([ - SplunkAOMetrics.correctness, - "completeness", - Metric(name="domain_relevance", version=3), - LocalMetricConfig(name="custom_metric", scorer_fn=custom_scorer) - ]) - - # Process local metrics if any - for local_metric in local_metrics: - logger.info(f"Need to process local metric: {local_metric.name}") - ``` - - Notes - ----- - **Requirements:** - - - The LogStream instance must have valid `id` and `project_id` attributes - - These are automatically set when retrieving LogStream objects via LogStreams methods - - **Recommended Usage:** - - - Use this method when you already have a LogStream object - - More intuitive than specifying project/log stream names again - - Cleaner object-oriented design pattern - """ - if not hasattr(self, "id") or not hasattr(self, "project_id"): - raise ValueError("Log stream must have id and project_id to enable metrics") - - _, local_metrics = create_metric_configs(self.project_id, self.id, metrics) - return local_metrics - - -class LogStreams: - config: SplunkAOConfig - - def __init__(self) -> None: - self.config = SplunkAOConfig.get() - - @overload - def list( - self, *, project_id: str, limit: Unset | int = 100, starting_token: Unset | int = 0 - ) -> builtins.list[LogStream]: ... - - @overload - def list( - self, *, project_name: str, limit: Unset | int = 100, starting_token: Unset | int = 0 - ) -> builtins.list[LogStream]: ... - - def list( - self, - *, - project_id: str | None = None, - project_name: str | None = None, - limit: Unset | int = 100, - starting_token: Unset | int = 0, - ) -> builtins.list[LogStream]: - """ - Lists log streams. Exactly one of `project_id` or `project_name` must be provided. - - Returns a single page of results. Use `starting_token` (from - `next_starting_token` on a prior response) to fetch subsequent pages. - - Parameters - ---------- - project_id : Optional[str], optional - The ID of the project to list log streams for. - project_name : Optional[str], optional - The name of the project to list log streams for. - limit : Union[Unset, int], optional - The maximum number of log streams to return per page. Defaults to 100. - starting_token : Union[Unset, int], optional - The pagination token to start from. Defaults to 0 (first page). - - Returns - ------- - builtins.list[LogStream] - A page of log streams. - - Raises - ------ - ValueError - If neither or both `project_id` and `project_name` are provided, - if the named project is not found, if the server returns a - validation error, or if the response is unexpectedly empty. - errors.UnexpectedStatus - If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. - httpx.TimeoutException - If the request takes longer than Client.timeout. - """ - if (project_id is None) == (project_name is None): - raise ValueError("Exactly one of 'project_id' or 'project_name' must be provided") - - if not project_id: - project = Projects().get(name=project_name) - if not project: - raise ValueError(f"Project {project_name} not found") - project_id = project.id - - response = list_log_streams_paginated_projects_project_id_log_streams_paginated_get.sync( - client=self.config.api_client, project_id=project_id, limit=limit, starting_token=starting_token - ) - - if isinstance(response, HTTPValidationError): - raise ValueError(f"Failed to list log streams: {response.detail}") - if response is None: - raise ValueError("Unexpected empty response while listing log streams") - - return [LogStream(log_stream=log_stream) for log_stream in response.log_streams] - - # Page size used by `_list_all`. Larger than the default `list()` page size so - # full scans (name-based `get`, oldest-stream fallback) issue fewer round trips. - # The OpenAPI spec does not declare a server-side maximum, so this value is a - # heuristic. If the server ever rejects it (e.g. 422), `_list_all` now raises - # `ValueError` instead of silently truncating, so the regression is loud. - _LIST_ALL_PAGE_SIZE = 500 - - def _list_all(self, *, project_id: str) -> builtins.list[LogStream]: - """Internal helper: paginate through every page and return all log streams. - - Used by callers that need a globally-complete view (e.g. name-based lookup, - oldest-stream fallback). Public callers should use `list()` with explicit - pagination instead. - - Raises ValueError on server validation failures or unexpected protocol - errors so mid-pagination failures aren't silently swallowed into a - truncated result. - """ - all_log_streams: builtins.list[LogStream] = [] - starting_token: int = 0 - seen_tokens: set[int] = {starting_token} - while True: - response = list_log_streams_paginated_projects_project_id_log_streams_paginated_get.sync( - client=self.config.api_client, - project_id=project_id, - starting_token=starting_token, - limit=self._LIST_ALL_PAGE_SIZE, - ) - if isinstance(response, HTTPValidationError): - raise ValueError(f"Failed to list log streams: {response.detail}") - if response is None: - raise ValueError("Unexpected empty response while paginating log streams") - - all_log_streams.extend(LogStream(log_stream=log_stream) for log_stream in response.log_streams) - - next_token = response.next_starting_token - if next_token is None or isinstance(next_token, Unset) or not response.paginated: - break - # Progress guard: stop if we've already seen this token. Catches both - # repeated and non-advancing tokens without assuming monotonic-integer - # ordering, so the loop stays safe if the server ever switches to - # opaque cursor tokens. - if not isinstance(next_token, int) or next_token in seen_tokens: - break - seen_tokens.add(next_token) - starting_token = next_token - - return all_log_streams - - @overload - def get(self, *, id: str, project_id: str | None = None, project_name: str | None = None) -> LogStream | None: ... - @overload - def get(self, *, name: str, project_id: str | None = None, project_name: str | None = None) -> LogStream | None: ... - def get( - self, - *, - id: str | None = None, - name: str | None = None, - project_id: str | None = None, - project_name: str | None = None, - ) -> LogStream | None: - """ - Retrieves a log stream by id or name. - - Parameters - ---------- - id : Optional[str], optional - The id of the log stream. Defaults to None. - name : Optional[str], optional - The name of the log stream. Defaults to None. - project_id : Optional[str], optional - The ID of the project. Defaults to None. - project_name : Optional[str], optional - The name of the project. Defaults to None. - - Returns - ------- - Optional[LogStream] - The log stream if found, None otherwise. - - Raises - ------ - ValueError - If neither or both `id` and `name` are provided, or if neither or both - `project_id` and `project_name` are provided. - errors.UnexpectedStatus - If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. - httpx.TimeoutException - If the request takes longer than Client.timeout. - """ - if (id is None) == (name is None): - raise ValueError("Exactly one of 'id' or 'name' must be provided") - - if (project_id is None) == (project_name is None): - raise ValueError("Exactly one of 'project_id' or 'project_name' must be provided") - - if not project_id: - project = Projects().get(name=project_name) - if not project: - raise ValueError(f"Project {project_name} not found") - project_id = project.id - - if id: - log_stream_response = get_log_stream_projects_project_id_log_streams_log_stream_id_get.sync( - project_id=project_id, log_stream_id=id, client=self.config.api_client - ) - if not log_stream_response: - return None - return LogStream(log_stream=log_stream_response) - - if name: - for log_stream in self._list_all(project_id=project_id): - if log_stream.name == name: - return log_stream - return None - - @overload - def create(self, name: str, *, project_id: str | None = None) -> LogStream: ... - @overload - def create(self, name: str, *, project_name: str) -> LogStream: ... - - def create(self, name: str, *, project_id: str | None = None, project_name: str | None = None) -> LogStream: - """ - Creates a new log stream. Exactly one of `project_id` or `project_name` must be provided. - - Parameters - ---------- - name : str - The name of the log stream. - project_id : Optional[str], optional - The ID of the project to create the log stream in. Defaults to None. - project_name : Optional[str], optional - The name of the project to create the log stream in. Defaults to None. - - Returns - ------- - LogStream - The created log stream. - - Raises - ------ - ValueError - If neither or both `project_id` and `project_name` are provided, or if the project is not found. - HTTPValidationError - If the server validation fails. - errors.UnexpectedStatus - If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. - httpx.TimeoutException - If the request takes longer than Client.timeout. - """ - if (project_id is None) == (project_name is None): - raise ValueError("Exactly one of 'project_id' or 'project_name' must be provided") - - if not project_id: - project = Projects().get(name=project_name) - if not project: - raise ValueError(f"Project {project_name} not found") - project_id = project.id - - body = LogStreamCreateRequest(name=name) - response = create_log_stream_projects_project_id_log_streams_post.sync( - project_id=project_id, client=self.config.api_client, body=body - ) - - if isinstance(response, HTTPValidationError): - raise response - - if not response: - raise ValueError("Unable to create log stream") - - return LogStream(log_stream=response) - - def enable_metrics( - self, - *, - log_stream_name: str | None = None, - project_name: str | None = None, - metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str], - ) -> builtins.list[LocalMetricConfig]: - """ - Enable metrics for a log stream by configuring scorers. - - The project name can be provided via the 'project_name' parameter or the - SPLUNK_AO_PROJECT environment variable. - - The log stream name can be provided via the 'log_stream_name' parameter or the - SPLUNK_AO_LOG_STREAM environment variable. - - Parameters - ---------- - log_stream_name : Optional[str], optional - The name of the log stream. Takes precedence over the SPLUNK_AO_LOG_STREAM environment variable. Defaults to None. - project_name : Optional[str], optional - The name of the project. Takes precedence over the SPLUNK_AO_PROJECT environment variable. Defaults to None. - metrics : builtins.list[Union[SplunkAOMetrics, Metric, LocalMetricConfig, str]] - List of metrics to enable. Can include: - - SplunkAOMetrics enum values (e.g., SplunkAOMetrics.correctness) - - Metric objects with name and optional version - - LocalMetricConfig objects for custom local metrics - - String names of built-in metrics - - Returns - ------- - tuple[builtins.list[ScorerConfig], builtins.list[LocalMetricConfig]] - A tuple containing the configured scorer configs and local metric configs. - - Raises - ------ - ValueError - If log stream or project cannot be found, or if metrics are unknown. - - Examples - -------- - ```python - # Enable built-in metrics with explicit parameters - from splunk_ao.log_streams import LogStreams - from splunk_ao.schema.metrics import SplunkAOMetrics - - log_streams = LogStreams() - scorer_configs, local_metrics = log_streams.enable_metrics( - log_stream_name="Production Logs", - project_name="My AI Project", - metrics=[ - SplunkAOMetrics.correctness, - SplunkAOMetrics.completeness, - "context_relevance", - ], - ) - - # Enable metrics using environment variables - # export SPLUNK_AO_LOG_STREAM="Production Logs" - # export SPLUNK_AO_PROJECT="My AI Project" - scorer_configs, local_metrics = log_streams.enable_metrics( - metrics=["correctness", "completeness"] - ) - - # Enable custom metrics with mixed parameters - from splunk_ao.schema.metrics import Metric, LocalMetricConfig - - def custom_scorer(trace_or_span): - return 0.85 # Custom scoring logic - - # export SPLUNK_AO_PROJECT="My AI Project" - scorer_configs, local_metrics = log_streams.enable_metrics( - log_stream_name="Production Logs", # Explicit log stream - # project_name from env var - metrics=[ - Metric(name="my_custom_metric", version=2), - LocalMetricConfig(name="local_scorer", scorer_fn=custom_scorer) - ] - ) - ``` - """ - # Apply environment variable fallbacks - project_name = project_name or _get_project_from_env() - log_stream_name = log_stream_name or _get_log_stream_from_env() - - # Get project using environment fallbacks - project_obj = Projects().get_with_env_fallbacks(name=project_name) - if not project_obj: - raise ValueError(f"Project '{project_name}' not found") - - # Get log stream - error out if not found - log_stream = self.get(name=log_stream_name, project_name=project_obj.name) - if not log_stream: - raise ValueError(f"Log stream '{log_stream_name}' not found in project '{project_obj.name}'") - - # Use the shared utility function directly - _, local_metrics = create_metric_configs(project_obj.id, log_stream.id, metrics) - return local_metrics - - -# -# Convenience methods -# - - -def get_log_stream( - *, name: str | None = None, project_id: str | None = None, project_name: str | None = None -) -> LogStream | None: - """ - Retrieves a log stream by name. Exactly one of `project_id` or `project_name` must be provided. - - Parameters - ---------- - name : Optional[str], optional - The name of the log stream. Defaults to None. - project_id : Optional[str], optional - The ID of the project. Defaults to None. - project_name : Optional[str], optional - The name of the project. Defaults to None. - - Returns - ------- - Optional[LogStream] - The log stream if found, None otherwise. - - Raises - ------ - ValueError - If neither or both `project_id` and `project_name` are provided. - errors.UnexpectedStatus - If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. - httpx.TimeoutException - If the request takes longer than Client.timeout. - """ - return LogStreams().get(name=name, project_id=project_id, project_name=project_name) - - -def list_log_streams( - *, - project_id: str | None = None, - project_name: str | None = None, - limit: Unset | int = 100, - starting_token: Unset | int = 0, -) -> builtins.list[LogStream]: - """ - Lists log streams. Exactly one of `project_id` or `project_name` must be provided. - - Returns a single page of results. Use `starting_token` (from - `next_starting_token` on a prior response) to fetch subsequent pages. - - Parameters - ---------- - project_id : str - The id of the project. - project_name : str - The name of the project. - limit : Union[Unset, int], optional - The maximum number of log streams to return per page. Defaults to 100. - starting_token : Union[Unset, int], optional - The pagination token to start from. Defaults to 0 (first page). - - Returns - ------- - builtins.list[LogStream] - A page of log streams. - - Raises - ------ - errors.UnexpectedStatus - If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. - httpx.TimeoutException - If the request takes longer than Client.timeout. - - """ - return LogStreams().list( - project_id=project_id, project_name=project_name, limit=limit, starting_token=starting_token - ) - - -def create_log_stream(name: str, project_id: str | None = None, project_name: str | None = None) -> LogStream: - """ - Creates a new log stream. Exactly one of `project_id` or `project_name` must be provided. - - Parameters - ---------- - name : str - The name of the log stream. - - Returns - ------- - LogStream - The created project. - - Raises - ------ - errors.UnexpectedStatus - If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. - httpx.TimeoutException - If the request takes longer than Client.timeout. - - """ - return LogStreams().create(name=name, project_id=project_id, project_name=project_name) - - -def enable_metrics( - *, - log_stream_name: str | None = None, - project_name: str | None = None, - metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str], -) -> builtins.list[LocalMetricConfig]: - """ - Enable metrics for a log stream with flexible parameter and environment variable support. - - This unified function supports both explicit parameters and environment variable fallbacks, - making it perfect for all use cases - from production CI/CD pipelines to development testing. - - **Flexible Usage Patterns:** - - **Environment-only**: Just pass metrics, names from env vars (production/CI) - - **Explicit parameters**: Specify project/log stream names directly (development) - - **Mixed approach**: Combine explicit params with environment fallbacks - - Environment Variables (Optional Fallbacks) - ------------------------------------------ - SPLUNK_AO_PROJECT : str - The name of the Galileo project (used when project_name not provided) - SPLUNK_AO_LOG_STREAM : str - The name of the log stream (used when log_stream_name not provided) - - Parameters - ---------- - log_stream_name : Optional[str], optional - The name of the log stream. Takes precedence over SPLUNK_AO_LOG_STREAM environment variable. - If None, will use SPLUNK_AO_LOG_STREAM env var. Defaults to None. - project_name : Optional[str], optional - The name of the project. Takes precedence over SPLUNK_AO_PROJECT environment variable. - If None, will use SPLUNK_AO_PROJECT env var. Defaults to None. - metrics : builtins.list[Union[SplunkAOMetrics, Metric, LocalMetricConfig, str]] - List of metrics to enable on the log stream. Can include: - - SplunkAOMetrics enum values (e.g., SplunkAOMetrics.correctness) - - Metric objects with name and optional version for custom metrics - - LocalMetricConfig objects for client-side custom scoring functions - - String names of built-in metrics (e.g., "correctness", "toxicity") - - Returns - ------- - builtins.list[LocalMetricConfig] - List of local metric configurations that must be computed client-side. - Server-side metrics are automatically registered with Galileo and don't - need to be returned since users don't interact with them. - - Raises - ------ - ValueError - If log stream or project cannot be found, or if any specified metrics are unknown. - errors.UnexpectedStatus - If the server returns an undocumented status code and Client.raise_on_unexpected_status is True. - httpx.TimeoutException - If the request takes longer than Client.timeout. - - Examples - -------- - ```python - # Enable built-in metrics with explicit parameters - from splunk_ao.log_streams import enable_metrics - from splunk_ao.schema.metrics import SplunkAOMetrics - - local_metrics = enable_metrics( - log_stream_name="Production Logs", - project_name="My AI Project", - metrics=[ - SplunkAOMetrics.correctness, - SplunkAOMetrics.completeness, - "context_relevance", - ], - ) - - # Enable metrics using environment variables only - # export SPLUNK_AO_LOG_STREAM="Production Logs" - # export SPLUNK_AO_PROJECT="My AI Project" - local_metrics = enable_metrics(metrics=["correctness", "completeness"]) - - # Enable custom and local metrics with environment variable fallbacks - from splunk_ao.schema.metrics import Metric, LocalMetricConfig - from galileo_core.schemas.logging.step import StepType - - def response_length_scorer(trace_or_span): - '''Custom metric that scores based on response length''' - if hasattr(trace_or_span, "output") and trace_or_span.output: - return min(len(trace_or_span.output) / 100.0, 1.0) # Normalize 0-1 - return 0.0 - - local_metrics = enable_metrics( - log_stream_name="Development Logs", - metrics=[ - SplunkAOMetrics.correctness, - "toxicity", - Metric(name="my_custom_metric", version=2), - LocalMetricConfig( - name="response_length", - scorer_fn=response_length_scorer, - scorable_types=[StepType.llm], - aggregatable_types=[StepType.trace], - ), - ], - ) - - # Process local metrics - for local_metric in local_metrics: - logger.info(f"Need to process local metric: {local_metric.name}") - ``` - """ - return LogStreams().enable_metrics(log_stream_name=log_stream_name, project_name=project_name, metrics=metrics) diff --git a/src/splunk_ao/logger/logger.py b/src/splunk_ao/logger/logger.py index 5879c938..78447221 100644 --- a/src/splunk_ao/logger/logger.py +++ b/src/splunk_ao/logger/logger.py @@ -38,7 +38,7 @@ from splunk_ao.constants import LoggerModeType from splunk_ao.constants.tracing import PARENT_ID_HEADER, TRACE_ID_HEADER from splunk_ao.exceptions import SplunkAOLoggerException -from splunk_ao.log_streams import LogStreams +from splunk_ao.agent_streams import AgentStreams from splunk_ao.logger.control import ControlAppliesTo, ControlCheckStage, ControlResult from splunk_ao.logger.task_handler import ThreadPoolTaskHandler from splunk_ao.projects import Projects @@ -383,7 +383,7 @@ def _init_project(self) -> None: @nop_sync def _init_log_stream(self) -> None: """Initializes the log stream ID.""" - log_streams_client = LogStreams() + log_streams_client = AgentStreams() log_stream_obj = log_streams_client.get(name=self.log_stream_name, project_id=self.project_id) if log_stream_obj is None: # Create log stream if it doesn't exist diff --git a/src/splunk_ao/metric.py b/src/splunk_ao/metric.py deleted file mode 100644 index 3825e362..00000000 --- a/src/splunk_ao/metric.py +++ /dev/null @@ -1,1317 +0,0 @@ -from __future__ import annotations - -import builtins -import json -import logging -import os -import time -from abc import ABC -from collections.abc import Callable -from datetime import datetime -from typing import TYPE_CHECKING, Any - -if TYPE_CHECKING: - from splunk_ao.model import Model - -from galileo_core.schemas.logging.span import Span -from galileo_core.schemas.logging.step import StepType -from galileo_core.schemas.logging.trace import Trace -from galileo_core.schemas.shared.metric import MetricValueType -from splunk_ao.config import SplunkAOConfig -from splunk_ao.configuration import Configuration -from splunk_ao.metrics import Metrics -from splunk_ao.resources.api.data import ( - create_code_scorer_version_scorers_scorer_id_version_code_post, - create_scorers_post, - get_validate_code_scorer_task_result_scorers_code_validate_task_id_get, - update_scorers_scorer_id_patch, - validate_code_scorer_scorers_code_validate_post, -) -from splunk_ao.resources.models import ( - BodyCreateCodeScorerVersionScorersScorerIdVersionCodePost, - BodyValidateCodeScorerScorersCodeValidatePost, - CreateScorerRequest, - HTTPValidationError, - OutputTypeEnum, - ScorerTypes, - TaskResultStatus, - UpdateScorerRequest, -) -from splunk_ao.resources.models.invalid_result import InvalidResult -from splunk_ao.resources.types import UNSET, File, Unset -from splunk_ao.schema.metrics import LocalMetricConfig, SplunkAOMetrics -from splunk_ao.schema.metrics import Metric as LegacyMetric -from splunk_ao.scorers import Scorers -from splunk_ao.shared.base import StateManagementMixin, SyncState -from splunk_ao.shared.exceptions import APIError, ValidationError - -logger = logging.getLogger(__name__) - -# Code validation polling parameters are configurable via Configuration: -# - Configuration.code_validation_timeout (env: SPLUNK_AO_CODE_VALIDATION_TIMEOUT) - default: 60.0s -# - Configuration.code_validation_initial_delay (env: SPLUNK_AO_CODE_VALIDATION_INITIAL_DELAY) - default: 5.0s -# - Configuration.code_validation_max_delay (env: SPLUNK_AO_CODE_VALIDATION_MAX_DELAY) - default: 30.0s -# - Configuration.code_validation_backoff_multiplier (env: SPLUNK_AO_CODE_VALIDATION_BACKOFF_MULTIPLIER) - default: 1.5 - - -class BuiltInMetrics: - """ - Provides convenient access to built-in Galileo metrics (formerly "scorers"). - - Examples - -------- - from splunk_ao.metric import Metric - - # Access built-in metrics - Metric.metrics.correctness - Metric.metrics.completeness - Metric.metrics.toxicity - """ - - def __getattr__(self, name: str) -> SplunkAOMetrics: - """Allow attribute-style access to built-in metrics.""" - # Try to find the metric by name (enum names match UI-visible names) - for scorer in SplunkAOMetrics: - if scorer.name == name: - return scorer - raise AttributeError(f"Built-in metric '{name}' not found. Available: {[s.name for s in SplunkAOMetrics]}") - - def __dir__(self) -> list[str]: - """Return list of available metric names for autocomplete.""" - return [scorer.name for scorer in SplunkAOMetrics] - - -# Backwards-compatible alias -BuiltInScorers = BuiltInMetrics - - -class Metric(StateManagementMixin, ABC): - """ - Base class for all Galileo metrics. - - This is an abstract base class that defines common attributes and methods - for all metric types. Use one of the concrete metric classes instead: - - - **SplunkAOMetric**: Built-in Galileo scorers (access via Metric.scorers) - - **LlmMetric**: Custom LLM-based metrics with prompt templates - - **LocalMetric**: Local function-based metrics - - **CodeMetric**: Code-based metrics (future support) - - Common Attributes - ----------------- - id (str | None): The unique metric identifier (UUID). - name (str): The metric name. - scorer_type (ScorerTypes | None): The type of scorer. - description (str): Description of the metric. - tags (list[str]): Tags associated with the metric. - created_at (datetime | None): When the metric was created. - updated_at (datetime | None): When the metric was last updated. - version (int | None): Metric version number. - - Class Attributes - ---------------- - metrics (BuiltInMetrics): Access built-in Galileo metrics. - - Examples - -------- - # 1. Use built-in Galileo scorers - from splunk_ao import Metric, SplunkAOMetric, LlmMetric, LocalMetric, LogStream - - log_stream = LogStream.get(name="my-stream", project_name="my-project") - log_stream.set_metrics([ - Metric.metrics.correctness, - Metric.metrics.completeness, - ]) - - # 2. Create custom LLM metric - llm_metric = LlmMetric( - name="response_quality", - prompt="Rate the quality...", - model="gpt-4o-mini", - judges=3, - ).create() - - # 3. Create local function-based metric - def my_scorer(trace_or_span): - return 0.5 - - local_metric = LocalMetric( - name="response_length", - scorer_fn=my_scorer, - ) - """ - - # Class attribute for built-in metrics (preferred name) - metrics = BuiltInMetrics() - - # Backwards-compatible property for legacy name - scorers = metrics - - # Type annotations for common instance attributes - id: str | None - name: str - scorer_type: ScorerTypes | None - description: str - tags: list[str] - created_at: datetime | None - updated_at: datetime | None - version: int | None - - # Scorer defaults - available for LLM and built-in Galileo metrics - # These are returned by the API in the ScorerDefaults object - model: str | None - judges: int | None - cot_enabled: bool | None - - def __init__( - self, name: str, *, description: str = "", tags: list[str] | None = None, version: int | None = None - ) -> None: - """ - Initialize a base Metric instance with common attributes. - - Args: - name: The name of the metric. - description: Description of the metric. - tags: Tags associated with the metric. - version: Specific version to reference (for existing metrics). - """ - super().__init__() - self.name = name - self.description = description - self.tags = tags if tags is not None else [] - self.version = version - self.id = None - self.created_at = None - self.updated_at = None - self.scorer_type = None - - # Initialize scorer defaults (populated from API for LLM and Galileo metrics) - self.model = None - self.judges = None - self.cot_enabled = None - - self._set_state(SyncState.LOCAL_ONLY) - - @staticmethod - def _parse_output_type( - output_type: str | OutputTypeEnum | None, default: OutputTypeEnum | None = None - ) -> OutputTypeEnum | None: - """Map a string or OutputTypeEnum to OutputTypeEnum, with an optional fallback default.""" - if output_type is None: - return default - if isinstance(output_type, OutputTypeEnum): - return output_type - _map = { - "percentage": OutputTypeEnum.PERCENTAGE, - "boolean": OutputTypeEnum.BOOLEAN, - "categorical": OutputTypeEnum.CATEGORICAL, - "count": OutputTypeEnum.COUNT, - "discrete": OutputTypeEnum.DISCRETE, - "freeform": OutputTypeEnum.FREEFORM, - "multilabel": OutputTypeEnum.MULTILABEL, - } - return _map.get(output_type.lower(), default) - - @classmethod - def _create_metric_from_type(cls, scorer_type: ScorerTypes) -> Metric: - """ - Create the appropriate Metric subclass instance based on scorer_type. - - This is a factory method that centralizes the logic for instantiating - the correct metric subclass based on the scorer type returned from the API. - - Args: - scorer_type: The scorer type from the API response. - - Returns - ------- - Metric: An uninitialized instance of the appropriate subclass - (LlmMetric, CodeMetric, or SplunkAOMetric). - - Examples - -------- - instance = Metric._create_metric_from_type(ScorerTypes.LLM) - # Returns: LlmMetric instance - """ - if scorer_type == ScorerTypes.LLM: - return LlmMetric.__new__(LlmMetric) - if scorer_type == ScorerTypes.CODE: - return CodeMetric.__new__(CodeMetric) - # Default to SplunkAOMetric for built-in scorers (LUNA, PRESET, etc.) - return SplunkAOMetric.__new__(SplunkAOMetric) - - @classmethod - def get(cls, *, id: str | None = None, name: str | None = None) -> Metric | None: - """ - Get an existing metric by ID or name. - - Returns the appropriate subclass instance based on scorer_type. - - Args: - id: The metric ID (UUID). - name: The metric name. - - Returns - ------- - Optional[Metric]: The metric if found (SplunkAOMetric, LlmMetric, or CodeMetric), None otherwise. - - Raises - ------ - ValidationError: If neither or both id and name are provided. - - Examples - -------- - # Get by name - returns appropriate subclass - metric = Metric.get(name="factuality-checker") - - # Get by ID - metric = Metric.get(id="abc-123-def") - """ - if id is not None and name is not None: - raise ValidationError("Cannot specify both id and name") - if id is None and name is None: - raise ValidationError("Must specify either id or name") - - scorers_service = Scorers() - - if name is not None: - scorers = scorers_service.list(name=name) - if not scorers: - return None - retrieved_scorer = next((s for s in scorers if s.name == name), None) - if retrieved_scorer is None: - return None - else: - assert id is not None - scorers = scorers_service.list() - retrieved_scorer = next((s for s in scorers if s.id == id), None) - if retrieved_scorer is None: - return None - - # Create appropriate subclass instance based on scorer_type - instance = cls._create_metric_from_type(retrieved_scorer.scorer_type) - StateManagementMixin.__init__(instance) - instance._populate_from_scorer_response(retrieved_scorer) - instance._set_state(SyncState.SYNCED) - return instance - - @classmethod - def list( - cls, *, name_filter: str | None = None, scorer_types: list[ScorerTypes] | None = None - ) -> builtins.list[Metric]: - """ - List metrics with optional filtering. - - Returns appropriate subclass instances based on scorer_type. - - Args: - name_filter: Filter metrics by exact name match. - scorer_types: Filter by scorer types. - - Returns - ------- - list[Metric]: List of metrics matching the criteria (with appropriate subclass types). - - Examples - -------- - # List all metrics - metrics = Metric.list() - - # List LLM metrics only - metrics = Metric.list(scorer_types=[ScorerTypes.LLM]) - - # List by name - metrics = Metric.list(name_filter="factuality") - """ - logger.debug(f"Metric.list: name_filter='{name_filter}' types={scorer_types} - started") - scorers_service = Scorers() - retrieved_scorers = scorers_service.list(name=name_filter, types=scorer_types) - logger.debug(f"Metric.list: found {len(retrieved_scorers)} metrics - completed") - - result: builtins.list[Metric] = [] - for retrieved_scorer in retrieved_scorers: - # Create appropriate subclass instance based on scorer_type - instance = cls._create_metric_from_type(retrieved_scorer.scorer_type) - StateManagementMixin.__init__(instance) - instance._populate_from_scorer_response(retrieved_scorer) - instance._set_state(SyncState.SYNCED) - result.append(instance) - - return result - - @classmethod - def delete_by_name(cls, name: str) -> None: - """ - Delete a metric by name without retrieving it first. - - This is more efficient than calling `Metric.get(name=...).delete()` - when you only need to delete and don't need the metric object. - - Args: - name: The name of the metric to delete. - - Raises - ------ - ValueError: If no metric with the given name exists. - - Examples - -------- - # Delete without retrieving first - Metric.delete_by_name("old-metric") - - # Alternative (less efficient) - metric = Metric.get(name="old-metric") - metric.delete() - """ - logger.info(f"Metric.delete_by_name: name='{name}' - started") - try: - metrics_service = Metrics() - metrics_service.delete_metric(name=name) - logger.info(f"Metric.delete_by_name: name='{name}' - completed") - except Exception as e: - logger.error(f"Metric.delete_by_name: name='{name}' - failed: {e}") - raise - - def _populate_from_scorer_response(self, scorer_response: Any) -> None: - """Populate instance attributes from a ScorerResponse object.""" - # Pre-compute optional common attributes - description = ( - "" - if isinstance(scorer_response.description, Unset) or scorer_response.description is None - else scorer_response.description - ) - created_at = None if isinstance(scorer_response.created_at, Unset) else scorer_response.created_at - updated_at = None if isinstance(scorer_response.updated_at, Unset) else scorer_response.updated_at - - # Extract defaults - available for LLM and built-in Galileo metrics - # These are returned by the API for preset scorers too - if not isinstance(scorer_response.defaults, Unset) and scorer_response.defaults is not None: - model = scorer_response.defaults.model_name if hasattr(scorer_response.defaults, "model_name") else None - judges = scorer_response.defaults.num_judges if hasattr(scorer_response.defaults, "num_judges") else None - cot_enabled = ( - scorer_response.defaults.cot_enabled if hasattr(scorer_response.defaults, "cot_enabled") else None - ) - else: - model = None - judges = None - cot_enabled = None - - self._sync_attrs( - id=scorer_response.id, - name=scorer_response.name, - scorer_type=scorer_response.scorer_type, - tags=scorer_response.tags, - version=None, - description=description, - created_at=created_at, - updated_at=updated_at, - model=model, - judges=judges, - cot_enabled=cot_enabled, - ) - - # LLM-specific attributes (only set if this is an LlmMetric) - if isinstance(self, LlmMetric): - output_type = None if isinstance(scorer_response.output_type, Unset) else scorer_response.output_type - prompt = None if isinstance(scorer_response.user_prompt, Unset) else scorer_response.user_prompt - - # Extract scoreable node types - if not isinstance(scorer_response.scoreable_node_types, Unset) and scorer_response.scoreable_node_types: - try: - node_level = StepType(scorer_response.scoreable_node_types[0]) - except (ValueError, IndexError): - node_level = None - else: - node_level = None - - ground_truth = ( - False - if isinstance(scorer_response.ground_truth, Unset) or scorer_response.ground_truth is None - else scorer_response.ground_truth - ) - - self._sync_attrs(output_type=output_type, prompt=prompt, node_level=node_level, ground_truth=ground_truth) - - # Code-specific attributes (only set if this is a CodeMetric) - if isinstance(self, CodeMetric): - # Extract scoreable node types - if not isinstance(scorer_response.scoreable_node_types, Unset) and scorer_response.scoreable_node_types: - try: - code_node_level = StepType(scorer_response.scoreable_node_types[0]) - except (ValueError, IndexError): - code_node_level = None - else: - code_node_level = None - - code_output_type = None if isinstance(scorer_response.output_type, Unset) else scorer_response.output_type - - self._sync_attrs(node_level=code_node_level, output_type=code_output_type) - - def update(self, **kwargs: Any) -> Metric: - """ - Update this metric's properties on the API. - - Only ``name``, ``description``, and ``tags`` can be updated via this method. - On success the instance is updated with the API response and returned in SYNCED state. - - Parameters - ---------- - **kwargs : Any - Fields to update. Supported keys: ``name``, ``description``, ``tags``. - - Returns - ------- - Metric: This metric instance with updated attributes from the API. - - Raises - ------ - ValidationError: If this is a local metric. - ValueError: If the metric ID is not set, the metric is deleted, - or the metric is in FAILED_SYNC state. - ValueError: If any unsupported fields are passed. - APIError: If the API returns a validation error or empty response. - Exception: If the API call fails (state is set to FAILED_SYNC). - - Examples - -------- - metric = Metric.get(name="factuality-checker") - metric.update(name="new-name", description="Updated description") - assert metric.is_synced() - """ - if isinstance(self, LocalMetric): - raise ValidationError("Local metrics don't exist on the server and can't be updated.") - if self.id is None: - raise ValueError("Metric ID is not set. Cannot update a local-only metric.") - if self.sync_state == SyncState.DELETED: - raise ValueError("Cannot update a deleted metric.") - if self.sync_state == SyncState.FAILED_SYNC: - raise ValueError( - "Cannot update a metric in FAILED_SYNC state. " - "Call refresh() to re-sync from the API, then retry your changes." - ) - - valid_fields = {"name", "description", "tags"} - invalid_fields = set(kwargs) - valid_fields - if invalid_fields: - raise ValueError(f"Invalid update fields: {sorted(invalid_fields)!r}. Valid fields: {sorted(valid_fields)}") - - body = UpdateScorerRequest( - name=kwargs.get("name", UNSET), description=kwargs.get("description", UNSET), tags=kwargs.get("tags", UNSET) - ) - - logger.info(f"Metric.update: id='{self.id}' name='{self.name}' - started") - try: - config = SplunkAOConfig.get() - response = update_scorers_scorer_id_patch.sync(scorer_id=self.id, client=config.api_client, body=body) - except Exception as e: - self._set_state(SyncState.FAILED_SYNC, error=e) - logger.error(f"Metric.update: id='{self.id}' - failed: {e}") - raise - - if isinstance(response, HTTPValidationError): - raise APIError(f"Metric update validation error: {response.detail}") - if response is None: - raise APIError(f"Metric update returned empty response for id '{self.id}'") - - self._populate_from_scorer_response(response) - self._set_state(SyncState.SYNCED) - logger.info(f"Metric.update: id='{self.id}' - completed") - return self - - def delete(self) -> None: - """ - Delete this metric. - - Only works for server-side metrics. Local metrics don't need deletion. - - Raises - ------ - ValidationError: If this is a local metric. - ValueError: If the metric is not synced. - - Examples - -------- - metric = Metric.get(name="factuality-checker") - metric.delete() - """ - if isinstance(self, LocalMetric): - raise ValidationError("Local metrics don't exist on the server and can't be deleted.") - - if self.id is None: - raise ValueError("Metric ID is not set. Cannot delete a local-only metric.") - - try: - logger.info(f"Metric.delete: id='{self.id}' name='{self.name}' - started") - metrics_service = Metrics() - metrics_service.delete_metric(name=self.name) - self._set_state(SyncState.DELETED) - logger.info(f"Metric.delete: id='{self.id}' - completed") - except Exception as e: - self._set_state(SyncState.FAILED_SYNC, error=e) - logger.error(f"Metric.delete: id='{self.id}' - failed: {e}") - raise - - def refresh(self) -> None: - """ - Refresh this metric's state from the API. - - Updates all attributes with the latest values from the remote API. - - Raises - ------ - ValidationError: If this is a local metric. - ValueError: If the metric is not synced. - Exception: If the API call fails or the metric no longer exists. - - Examples - -------- - metric.refresh() - assert metric.is_synced() - """ - if isinstance(self, LocalMetric): - raise ValidationError("Local metrics don't exist on the server and can't be refreshed.") - - if self.id is None: - raise ValueError("Metric ID is not set. Cannot refresh a local-only metric.") - - try: - logger.debug(f"Metric.refresh: id='{self.id}' - started") - scorers_service = Scorers() - scorers = scorers_service.list() - retrieved_scorer = next((s for s in scorers if s.id == self.id), None) - - if retrieved_scorer is None: - raise ValueError(f"Metric with id '{self.id}' no longer exists") - - self._populate_from_scorer_response(retrieved_scorer) - self._set_state(SyncState.SYNCED) - logger.debug(f"Metric.refresh: id='{self.id}' - completed") - except Exception as e: - self._set_state(SyncState.FAILED_SYNC, error=e) - logger.error(f"Metric.refresh: id='{self.id}' - failed: {e}") - raise - - def to_legacy_metric(self) -> LegacyMetric: - """ - Convert to legacy splunk_ao.schema.metrics.Metric format. - - This enables backward compatibility with existing code that uses - the legacy Metric class. - - Returns - ------- - LegacyMetric: Legacy metric object with name and version. - - Examples - -------- - metric = Metric.get(name="my-metric") - legacy = metric.to_legacy_metric() - # Use with existing APIs - """ - return LegacyMetric(name=self.name, version=self.version) - - def __str__(self) -> str: - """String representation of the metric.""" - type_name = self.__class__.__name__ - scorer_type_str = self.scorer_type.value if self.scorer_type else "unknown" - return f"{type_name}(name='{self.name}', id='{self.id}', scorer_type='{scorer_type_str}')" - - def __repr__(self) -> str: - """Detailed string representation of the metric.""" - type_name = self.__class__.__name__ - return f"{type_name}(name='{self.name}', id='{self.id}')" - - -# ============================================================================ -# Concrete Metric Types -# ============================================================================ - - -class LlmMetric(Metric): - """ - LLM-based metric with custom prompt templates. - - This metric type allows you to create custom metrics evaluated by an LLM - judge using a prompt template. - - Attributes - ---------- - prompt (str | None): Prompt template for the LLM scorer. - model (str | None): Model name/alias to use for scoring (stored as string). - judges (int | None): Number of judges to use for scoring. - cot_enabled (bool | None): Whether chain-of-thought is enabled. - node_level (StepType | None): Node level for the metric. - output_type (OutputTypeEnum | None): Output type for the metric. - - Configuration - ------------- - Default values for `model` and `judges` can be configured via: - - Configuration.default_scorer_model (env: SPLUNK_AO_DEFAULT_SCORER_MODEL) - - Configuration.default_scorer_judges (env: SPLUNK_AO_DEFAULT_SCORER_JUDGES) - - Examples - -------- - # Create custom LLM metric with string model name - metric = LlmMetric( - name="response_quality", - prompt=''' - Rate the quality of this response on a scale of 1-10. - - Question: {input} - Answer: {output} - - Return only the numerical score (1-10). - ''', - model="gpt-4o-mini", # String model name - judges=3, - node_level=StepType.llm, - description="Rates response quality", - tags=["quality", "custom"], - output_type=OutputTypeEnum.PERCENTAGE, - cot_enabled=True, - ).create() - - # Or use a Model object from Integration - from splunk_ao.integration import Integration - gpt_model = Integration.openai.get_model(alias="gpt-4o-mini") - metric = LlmMetric( - name="response_quality", - prompt="Rate quality 1-10: {input} -> {output}", - model=gpt_model, # Model object - judges=3, - ).create() - """ - - # Type annotations for LLM-specific attributes - prompt: str | None - model: str | None - judges: int | None - cot_enabled: bool | None - node_level: StepType | None - output_type: OutputTypeEnum | None - ground_truth: bool - - def __init__( - self, - name: str, - *, - # LLM metric parameters (improved API) - prompt: str | None = None, - model: Model | str | None = None, - judges: int | None = None, - # Backward compatibility aliases - user_prompt: str | None = None, - model_name: str | None = None, - num_judges: int | None = None, - # LLM-specific parameters - node_level: StepType | None = None, - cot_enabled: bool | None = None, - output_type: str | OutputTypeEnum | None = None, - ground_truth: bool = False, - # Common parameters - description: str = "", - tags: list[str] | None = None, - version: int | None = None, - ) -> None: - """ - Initialize an LLM metric. - - Args: - name: The name of the metric. - prompt: Prompt template for LLM scorers (preferred over user_prompt). - model: Model object or model name string to use (preferred over model_name). - Defaults to Configuration.default_scorer_model. - judges: Number of judges (preferred over num_judges). Defaults to Configuration.default_scorer_judges. - user_prompt: [Deprecated] Use 'prompt' instead. - model_name: [Deprecated] Use 'model' instead. - num_judges: [Deprecated] Use 'judges' instead. - node_level: Node level for the metric. Defaults to StepType.llm. - cot_enabled: Whether chain-of-thought is enabled. Defaults to True. - output_type: Output type ("percentage", "boolean", etc.). - ground_truth: Whether the scorer requires ground truth (``reference_output``) from the dataset. - When True, the judge LLM receives the row's ground-truth value in its prompt. - description: Description of the metric. - tags: Tags associated with the metric. - version: Specific version to reference (for existing metrics). - - Raises - ------ - ValidationError: If prompt is not provided. - """ - super().__init__(name=name, description=description, tags=tags, version=version) - - # Handle parameter aliases (new names preferred) - final_prompt = prompt or user_prompt - - # Handle model parameter - extract alias from Model object if needed - if model is not None: - # Local import to avoid circular dependency - from splunk_ao.model import Model - - final_model = model.alias if isinstance(model, Model) else model - else: - final_model = model_name or Configuration.default_scorer_model - - final_judges = ( - judges - if judges is not None - else (num_judges if num_judges is not None else Configuration.default_scorer_judges) - ) - - if final_prompt is None: - raise ValidationError("'prompt' (or 'user_prompt') must be provided for LLM-based metrics.") - - # Initialize LLM-specific attributes - self.prompt = final_prompt - self.model = final_model # Now always a string (alias) - self.judges = final_judges - self.node_level = node_level or StepType.llm - self.cot_enabled = cot_enabled if cot_enabled is not None else True - - # Handle output_type (accept string or enum) - if isinstance(output_type, str): - self.output_type = Metric._parse_output_type(output_type, default=OutputTypeEnum.PERCENTAGE) - else: - self.output_type = output_type or OutputTypeEnum.BOOLEAN - - self.ground_truth = ground_truth - - self.scorer_type = ScorerTypes.LLM - - def create(self) -> LlmMetric: - """ - Persist this LLM metric to the API. - - Returns - ------- - LlmMetric: This metric instance with updated attributes from the API. - - Raises - ------ - ValidationError: If configuration is invalid. - Exception: If the API call fails. - - Examples - -------- - metric = LlmMetric( - name="quality_check", - prompt="Rate the quality...", - model="gpt-4o-mini" - ).create() - assert metric.is_synced() - """ - try: - logger.info(f"LlmMetric.create: name='{self.name}' - started") - - metrics_service = Metrics() - created_version = metrics_service.create_custom_llm_metric( - name=self.name, - user_prompt=self.prompt or "", - node_level=self.node_level if self.node_level is not None else StepType.llm, - cot_enabled=self.cot_enabled if self.cot_enabled is not None else True, - model_name=self.model if self.model is not None else Configuration.default_scorer_model, - num_judges=self.judges if self.judges is not None else Configuration.default_scorer_judges, - description=self.description, - tags=self.tags, - output_type=self.output_type - if isinstance(self.output_type, OutputTypeEnum) - else OutputTypeEnum.BOOLEAN, - ground_truth=self.ground_truth, - ) - - # Update attributes from response without triggering dirty-tracking - self._sync_attrs( - id=str(created_version.scorer_id), - created_at=created_version.created_at, - updated_at=created_version.updated_at, - ) - - # Refresh to get full scorer details - self.refresh() - - logger.info(f"LlmMetric.create: id='{self.id}' - completed") - return self - except ValidationError: - raise - except Exception as e: - self._set_state(SyncState.FAILED_SYNC, error=e) - logger.error(f"LlmMetric.create: name='{self.name}' - failed: {e}") - raise - - def __repr__(self) -> str: - """Detailed string representation of the metric.""" - return f"LlmMetric(name='{self.name}', id='{self.id}', model='{self.model}', judges={self.judges})" - - -class CodeMetric(Metric): - r""" - Code-based metric. - - This metric type is for code-based scorers that execute custom code - to evaluate traces/spans. - - Attributes - ---------- - node_level (StepType | None): Node level for the metric. - code (str | None): The Python code for the scorer. - output_type (OutputTypeEnum | None): Output type for the metric. - - Examples - -------- - # Get existing code metric - metric = Metric.get(name="my-code-metric") - assert isinstance(metric, CodeMetric) - - # Create code metric with inline code - metric = CodeMetric( - name="custom_code_scorer", - code="def scorer_fn(step_object):\\n return 1.0", - description="Custom code-based scorer", - tags=["custom", "code"], - node_level=StepType.llm, - output_type=OutputTypeEnum.PERCENTAGE, - ).create() - - # Load code from file - metric = CodeMetric( - name="custom_code_scorer", - node_level=StepType.llm, - ).load_code("./scorers/my_scorer.py").create() - """ - - # Type annotations for code-specific attributes - node_level: StepType | None - code: str | None - output_type: OutputTypeEnum | None - required_metrics: list[str] | None - - def __init__( - self, - name: str, - *, - code: str | None = None, - node_level: StepType | None = None, - output_type: str | OutputTypeEnum | None = None, - required_metrics: list[str] | None = None, - description: str = "", - tags: list[str] | None = None, - version: int | None = None, - ) -> None: - """ - Initialize a Code metric. - - Args: - name: The name of the metric. - code: The Python code for the scorer (optional, can be set later or loaded from file). - node_level: Node level for the metric. Defaults to StepType.llm. - output_type: Output type for the metric ("percentage", "boolean", "categorical", - "count", "discrete"). Accepts string or OutputTypeEnum. - required_metrics: List of metric names that this code metric depends on. - description: Description of the metric. - tags: Tags associated with the metric. - version: Specific version to reference (for existing metrics). - """ - super().__init__(name=name, description=description, tags=tags, version=version) - - self.code = code - self.node_level = node_level or StepType.llm - self.required_metrics = required_metrics - self.scorer_type = ScorerTypes.CODE - - self.output_type = Metric._parse_output_type(output_type) - - def load_code(self, code_file_path: str) -> CodeMetric: - """ - Load code from a file into this metric instance. - - Args: - code_file_path: Path to the Python file containing the scorer code. - - Returns - ------- - CodeMetric: This metric instance with code loaded from the file (for chaining). - - Raises - ------ - ValidationError: If the code file doesn't exist or can't be read. - - Examples - -------- - # Load code from file - metric = CodeMetric( - name="custom_code_scorer", - node_level=StepType.llm, - ).load_code("./scorers/my_scorer.py").create() - """ - if not os.path.isfile(code_file_path): - raise ValidationError(f"Code file not found: {code_file_path}") - - try: - with open(code_file_path, encoding="utf-8") as f: - self.code = f.read() - except Exception as e: - raise ValidationError(f"Failed to read code file {code_file_path}: {e}") - - return self - - def _validate_code(self, config: SplunkAOConfig) -> str: - """ - Validate the code by submitting it to the validation endpoint and polling for results. - - Args: - config: The Galileo configuration with API client. - - Returns - ------- - str: The validation result as a JSON string to pass to create_code_scorer_version. - - Raises - ------ - ValidationError: If validation fails or the code is invalid. - ValueError: If the API returns an unexpected response. - """ - assert self.code is not None - assert self.node_level is not None - - # Step 1: Submit the code for validation - code_bytes = self.code.encode("utf-8") - code_file = File(payload=code_bytes, file_name="scorer.py") - validate_body = BodyValidateCodeScorerScorersCodeValidatePost( - file=code_file, scoreable_node_types=[self.node_level.value], required_scorers=self.required_metrics - ) - - validate_response = validate_code_scorer_scorers_code_validate_post.sync( - client=config.api_client, body=validate_body - ) - - if validate_response is None: - logger.debug("CodeMetric._validate_code: No response from validate_code_scorer") - raise ValueError("Failed to validate code: No response from API") - - task_id = validate_response.task_id - logger.debug(f"CodeMetric._validate_code: task_id='{task_id}' - validation started") - - # Step 2: Poll for validation result with time-based timeout - start_time_seconds = time.time() - attempt = 0 - while True: - elapsed_seconds = time.time() - start_time_seconds - timeout_seconds = Configuration.code_validation_timeout - if elapsed_seconds >= timeout_seconds: - raise ValidationError(f"Code validation timed out after {timeout_seconds:.0f} seconds") - - task_result = get_validate_code_scorer_task_result_scorers_code_validate_task_id_get.sync( - task_id=task_id, client=config.api_client - ) - - if task_result is None: - logger.debug(f"CodeMetric._validate_code: No response for task_id='{task_id}'") - raise ValueError("Failed to get validation result: No response from API") - - if task_result.status == TaskResultStatus.COMPLETED: - logger.debug(f"CodeMetric._validate_code: task_id='{task_id}' - validation completed") - - # Extract and validate the result - result = task_result.result - - # Handle string result (already serialized) - if isinstance(result, str): - return result - - # Handle ValidateRegisteredScorerResult or similar objects with to_dict - if hasattr(result, "to_dict"): - # Check if it's an invalid result (has error_message in nested result) - if hasattr(result, "result") and isinstance(result.result, InvalidResult): - raise ValidationError(f"Code validation failed: {result.result.error_message}") - # Return the result as JSON string - return json.dumps(result.to_dict()) - - raise ValueError(f"Unexpected validation result type: {type(result)}") - - if task_result.status == TaskResultStatus.FAILED: - error_msg = "Code validation failed" - if isinstance(task_result.result, str): - error_msg = f"Code validation failed: {task_result.result}" - raise ValidationError(error_msg) - - if task_result.status == TaskResultStatus.PENDING: - # Calculate delay with exponential backoff - delay_seconds = min( - Configuration.code_validation_initial_delay - * (Configuration.code_validation_backoff_multiplier**attempt), - Configuration.code_validation_max_delay, - ) - - logger.debug( - f"CodeMetric._validate_code: task_id='{task_id}' - pending " - f"(elapsed: {elapsed_seconds:.1f}s/{timeout_seconds:.0f}s, next delay: {delay_seconds:.2f}s)" - ) - time.sleep(delay_seconds) - attempt += 1 - else: - raise ValueError(f"Unknown task status: {task_result.status}") - - def create(self) -> CodeMetric: - r""" - Persist this Code metric to the API. - - This method validates the code first by submitting it to the validation - endpoint, polling for the result, and then creating the scorer with the - validated result. - - Returns - ------- - CodeMetric: This metric instance with updated attributes from the API. - - Raises - ------ - ValidationError: If code is not set, validation fails, or configuration is invalid. - Exception: If the API call fails. - - Examples - -------- - # Create with inline code - metric = CodeMetric( - name="custom_code_scorer", - code="def scorer_fn(step_object):\\n return 1.0", - node_level=StepType.llm, - ).create() - assert metric.is_synced() - - # Create by loading from file - metric = CodeMetric( - name="custom_code_scorer", - node_level=StepType.llm, - ).load_code("./scorers/my_scorer.py").create() - assert metric.is_synced() - """ - # Validate that code is set - if self.code is None: - raise ValidationError( - "Code is not set. Either pass 'code' to __init__() or use CodeMetric.load_code() to load from a file." - ) - - try: - logger.info(f"CodeMetric.create: name='{self.name}' - started") - - config = SplunkAOConfig.get() - - # Ensure node_level is set (should always be set in __init__, but checking for type safety) - assert self.node_level is not None - - # Step 1: Validate the code and get validation result - logger.debug(f"CodeMetric.create: name='{self.name}' - validating code") - validation_result = self._validate_code(config) - logger.debug(f"CodeMetric.create: name='{self.name}' - code validated successfully") - - # Step 2: Create the scorer - scorer_request = CreateScorerRequest( - name=self.name, - scorer_type=ScorerTypes.CODE, - description=self.description, - tags=self.tags, - scoreable_node_types=[self.node_level.value], - output_type=self.output_type, - required_scorers=self.required_metrics, - ) - - scorer_response = create_scorers_post.sync(client=config.api_client, body=scorer_request) - - if scorer_response is None: - logger.debug("CodeMetric.create: No response from create_scorers_post") - raise ValueError("Failed to create code-based metric: No response from API") - - # Step 3: Create the code scorer version with file upload and validation result - # Convert the code string to bytes for file upload - code_bytes = self.code.encode("utf-8") - code_file = File(payload=code_bytes, file_name="scorer.py") - version_body = BodyCreateCodeScorerVersionScorersScorerIdVersionCodePost( - file=code_file, validation_result=validation_result - ) - - created_version = create_code_scorer_version_scorers_scorer_id_version_code_post.sync( - scorer_id=scorer_response.id, client=config.api_client, body=version_body - ) - - if created_version is None: - logger.debug( - "CodeMetric.create: No response from create_code_scorer_version_scorers_scorer_id_version_code_post" - ) - raise ValueError("Failed to create code-based metric: No response from API") - - # Update attributes from response - self.id = str(scorer_response.id) - self.created_at = scorer_response.created_at - self.updated_at = scorer_response.updated_at - - # Refresh to get full scorer details - self.refresh() - - logger.info(f"CodeMetric.create: id='{self.id}' - completed") - return self - except ValidationError: - raise - except Exception as e: - self._set_state(SyncState.FAILED_SYNC, error=e) - logger.error(f"CodeMetric.create: name='{self.name}' - failed: {e}") - raise - - def __repr__(self) -> str: - """Detailed string representation of the metric.""" - return f"CodeMetric(name='{self.name}', id='{self.id}'')" - - -class SplunkAOMetric(Metric): - """ - Built-in Galileo scorer metric. - - This metric type represents Galileo's built-in scorers like correctness, - completeness, toxicity, etc. Access these via `Metric.metrics`. - - Examples - -------- - # Access built-in scorers - from splunk_ao import Metric, LogStream - - log_stream = LogStream.get(name="my-stream", project_name="my-project") - log_stream.set_metrics([ - Metric.metrics.correctness, - Metric.metrics.completeness, - Metric.metrics.toxicity, - ]) - - # Or get by name - metric = Metric.get(name="correctness") - assert isinstance(metric, SplunkAOMetric) - """ - - def __init__( - self, name: str, *, description: str = "", tags: list[str] | None = None, version: int | None = None - ) -> None: - """ - Initialize a Galileo metric. - - Args: - name: The name of the metric. - description: Description of the metric. - tags: Tags associated with the metric. - version: Specific version to reference (for existing metrics). - """ - super().__init__(name=name, description=description, tags=tags, version=version) - # Galileo metrics can have various scorer types, set during population - - -class LocalMetric(Metric): - """ - Local function-based metric. - - This metric type uses a Python function to score traces/spans locally - without making API calls. Useful for simple, deterministic metrics. - - Attributes - ---------- - scorer_fn (Callable): Scoring function that takes a Trace or Span and returns either a - score, or a ``(score, metadata)`` tuple — when a tuple is returned, the metadata dict - is attached to the step under ``{name}_metadata`` for explainability. - scorable_types (list[StepType]): Types that can be scored. - aggregatable_types (list[StepType]): Types that can be aggregated. - - Examples - -------- - # Create local function-based metric - def response_length_scorer(trace_or_span): - if hasattr(trace_or_span, "output") and trace_or_span.output: - return min(len(trace_or_span.output) / 100.0, 1.0) - return 0.0 - - local_metric = LocalMetric( - name="response_length", - scorer_fn=response_length_scorer, - scorable_types=[StepType.llm], - aggregatable_types=[StepType.trace], - ) - - # Or return (score, metadata) for explainability - EXPECTED = ["relevance", "accuracy", "completeness"] - def keyword_coverage(trace_or_span): - text = getattr(trace_or_span, "output", "") or "" - matched = [k for k in EXPECTED if k in text] - return len(matched) / len(EXPECTED), { - "matched": matched, - "missing": [k for k in EXPECTED if k not in text], - } - - # Use with log stream - log_stream.set_metrics([local_metric]) - """ - - # Type annotations for local metric attributes - scorer_fn: Callable[[Trace | Span], MetricValueType | tuple[MetricValueType, dict[str, Any]]] - scorable_types: list[StepType] - aggregatable_types: list[StepType] - - def __init__( - self, - name: str, - *, - scorer_fn: Callable[[Trace | Span], MetricValueType | tuple[MetricValueType, dict[str, Any]]], - scorable_types: list[StepType] | None = None, - aggregatable_types: list[StepType] | None = None, - description: str = "", - tags: list[str] | None = None, - ) -> None: - """ - Initialize a local function-based metric. - - Args: - name: The name of the metric. - scorer_fn: Scoring function for the metric. May return either a bare score, or a - ``(score, metadata)`` tuple where ``metadata`` is a JSON-serializable dict - surfaced under ``{name}_metadata`` on the step. - scorable_types: Step types that can be scored. Defaults to [StepType.llm]. - aggregatable_types: Step types for aggregation. Defaults to [StepType.trace]. - description: Description of the metric. - tags: Tags associated with the metric. - - Raises - ------ - ValidationError: If scorer_fn is not provided. - """ - super().__init__(name=name, description=description, tags=tags) - - if scorer_fn is None: - raise ValidationError("'scorer_fn' must be provided for local metrics.") - - self.scorer_fn = scorer_fn - self.scorable_types = scorable_types or [StepType.llm] - self.aggregatable_types = aggregatable_types or [StepType.trace] - self.scorer_type = None # Local metrics don't have a scorer_type - - def to_local_metric_config(self) -> LocalMetricConfig: - """ - Convert to LocalMetricConfig format. - - Returns - ------- - LocalMetricConfig: Local metric configuration for use with the logger. - - Examples - -------- - def my_scorer(trace): - return 0.5 - - metric = LocalMetric(name="test", scorer_fn=my_scorer) - config = metric.to_local_metric_config() - """ - return LocalMetricConfig( - name=self.name, - scorer_fn=self.scorer_fn, - scorable_types=self.scorable_types, - aggregatable_types=self.aggregatable_types, - ) - - def __repr__(self) -> str: - """Detailed string representation of the metric.""" - # Handle callables that don't have __name__ (partials, lambdas, callable instances) - fn_name = getattr(self.scorer_fn, "__name__", f"<{type(self.scorer_fn).__name__}>") - return f"LocalMetric(name='{self.name}', scorer_fn={fn_name})" diff --git a/src/splunk_ao/metrics.py b/src/splunk_ao/metrics.py deleted file mode 100644 index 3aa864c7..00000000 --- a/src/splunk_ao/metrics.py +++ /dev/null @@ -1,259 +0,0 @@ -import datetime -import logging - -from galileo_core.schemas.logging.step import StepType -from splunk_ao.config import SplunkAOConfig -from splunk_ao.resources.api.data import ( - create_llm_scorer_version_scorers_scorer_id_version_llm_post, - create_scorers_post, - delete_scorer_scorers_scorer_id_delete, -) -from splunk_ao.resources.api.trace import query_metrics_projects_project_id_metrics_search_post -from splunk_ao.resources.models import ( - HTTPValidationError, - LogRecordsMetricsQueryRequest, - LogRecordsMetricsResponse, - ScorerTypes, -) -from splunk_ao.resources.models.base_scorer_version_response import BaseScorerVersionResponse -from splunk_ao.resources.models.create_llm_scorer_version_request import CreateLLMScorerVersionRequest -from splunk_ao.resources.models.create_scorer_request import CreateScorerRequest -from splunk_ao.resources.models.output_type_enum import OutputTypeEnum -from splunk_ao.resources.models.scorer_defaults import ScorerDefaults -from splunk_ao.scorers import Scorers -from splunk_ao.search import FilterType - -_logger = logging.getLogger(__name__) - - -class Metrics: - config: SplunkAOConfig - - def __init__(self) -> None: - self.config = SplunkAOConfig.get() - - def delete_metric(self, name: str) -> None: - scorers_to_delete = Scorers().list(name=name) - if not scorers_to_delete: - raise ValueError(f"Scorer with name {name} not found.") - - for scorer in scorers_to_delete: - response = delete_scorer_scorers_scorer_id_delete.sync(scorer_id=scorer.id, client=self.config.api_client) - - if isinstance(response, HTTPValidationError): - raise ValueError(response.detail) - if response is None: - raise ValueError("Failed to delete metric.") - - def create_custom_llm_metric( - self, - name: str, - user_prompt: str, - node_level: StepType = StepType.llm, - cot_enabled: bool = True, - model_name: str = "gpt-4.1-mini", - num_judges: int = 3, - description: str = "", - tags: list[str] | None = None, - output_type: OutputTypeEnum = OutputTypeEnum.BOOLEAN, - ground_truth: bool = False, - ) -> BaseScorerVersionResponse: - """ - Create a custom LLM metric. - - Parameters - ---------- - name: str - Name of the metric. - user_prompt: str - User prompt for the metric. - node_level: StepType - Node level for the metric. - cot_enabled: bool - Whether chain-of-thought is enabled. - model_name: str - Model name to use. - num_judges: int - Number of judges for the metric. - description: str - Description of the metric. - tags: List[str] - Tags associated with the metric. - output_type: OutputTypeEnum - Output type for the metric. - ground_truth: bool - Whether the scorer requires ground truth (``reference_output``) from the dataset. - When True, the judge LLM receives the row's ground-truth value in its prompt. - - Returns - ------- - BaseScorerVersionResponse - Response containing the created metric details. - """ - if tags is None: - tags = [] - create_scorer_request = CreateScorerRequest( - name=name, - scorer_type=ScorerTypes.LLM, - description=description, - tags=tags, - defaults=ScorerDefaults(model_name=model_name, num_judges=num_judges, cot_enabled=cot_enabled), - scoreable_node_types=[node_level], - output_type=output_type, - ground_truth=ground_truth, - ) - - scorer = create_scorers_post.sync(body=create_scorer_request, client=self.config.api_client) - - version_req = CreateLLMScorerVersionRequest(user_prompt=user_prompt) - version_resp = create_llm_scorer_version_scorers_scorer_id_version_llm_post.sync( - scorer_id=scorer.id, body=version_req, client=self.config.api_client - ) - - _logger.info("Created custom LLM metric: %s", name) - - return version_resp - - def query( - self, - project_id: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - experiment_id: str | None = None, - log_stream_id: str | None = None, - filters: list[FilterType] | None = None, - group_by: str | None = None, - interval: int = 5, - ) -> LogRecordsMetricsResponse: - body = LogRecordsMetricsQueryRequest( - start_time=start_time, - end_time=end_time, - experiment_id=experiment_id, - log_stream_id=log_stream_id, - filters=filters or [], - group_by=group_by, - interval=interval, - ) - - response = query_metrics_projects_project_id_metrics_search_post.sync( - client=self.config.api_client, project_id=str(project_id), body=body - ) - - if isinstance(response, HTTPValidationError): - raise ValueError(response.detail) - if response is None: - raise ValueError("Failed to query for metrics.") - - return response - - -# Public functions -def create_custom_llm_metric( - name: str, - user_prompt: str, - node_level: StepType = StepType.llm, - cot_enabled: bool = True, - model_name: str = "gpt-4.1-mini", - num_judges: int = 3, - description: str = "", - tags: list[str] | None = None, - output_type: OutputTypeEnum = OutputTypeEnum.BOOLEAN, - ground_truth: bool = False, -) -> BaseScorerVersionResponse: - """ - Create a custom LLM metric. - - Parameters - ---------- - name: str - Name of the metric. - user_prompt: str - User prompt for the metric. - node_level: StepType - Node level for the metric. - cot_enabled: bool - Whether chain-of-thought is enabled. - model_name: str - Model name to use. - num_judges: int - Number of judges for the metric. - description: str - Description of the metric. - tags: List[str] - Tags associated with the metric. - output_type: OutputTypeEnum - Output type for the metric. - ground_truth: bool - Whether the scorer requires ground truth (``reference_output``) from the dataset. - When True, the judge LLM receives the row's ground-truth value in its prompt. - - Returns - ------- - BaseScorerVersionResponse - Response containing the created metric details. - """ - if tags is None: - tags = [] - return Metrics().create_custom_llm_metric( - name, user_prompt, node_level, cot_enabled, model_name, num_judges, description, tags, output_type, ground_truth - ) - - -def get_metrics( - project_id: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - experiment_id: str | None = None, - log_stream_id: str | None = None, - filters: list[FilterType] | None = None, - group_by: str | None = None, - interval: int = 5, -) -> LogRecordsMetricsResponse: - """Queries for metrics in a project. - - Parameters - ---------- - project_id - The unique identifier of the project. - start_time - The start of the time range for the query. - end_time - The end of the time range for the query. - experiment_id - Filter records by a specific experiment ID. - log_stream_id - Filter records by a specific run ID. - filters - A list of filters to apply to the query. - group_by - The field to group the results by. - interval - The time interval for the query in seconds. - - Returns - ------- - LogRecordsMetricsResponse - A LogRecordsMetricsResponse object containing the query results, or None if the query fails. - """ - return Metrics().query( - project_id=project_id, - start_time=start_time, - end_time=end_time, - experiment_id=experiment_id, - log_stream_id=log_stream_id, - filters=filters, - group_by=group_by, - interval=interval, - ) - - -def delete_metric(name: str) -> None: - """ - Deletes a metric by its name. - - Parameters - ---------- - name - The name of the metric to delete. - """ - Metrics().delete_metric(name) diff --git a/src/splunk_ao/project.py b/src/splunk_ao/project.py index 9ae5e9e9..d06d6874 100644 --- a/src/splunk_ao/project.py +++ b/src/splunk_ao/project.py @@ -3,7 +3,7 @@ import builtins import logging from datetime import datetime -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any from splunk_ao.collaborator import Collaborator, CollaboratorRole from splunk_ao.config import SplunkAOConfig @@ -19,7 +19,6 @@ from splunk_ao.agent_stream import AgentStream from splunk_ao.dataset import Dataset from splunk_ao.experiment import Experiment - from splunk_ao.log_stream import LogStream from splunk_ao.prompt import Prompt logger = logging.getLogger(__name__) @@ -56,14 +55,14 @@ class Project(StateManagementMixin): projects = Project.list() # Create a log stream for the project - log_stream = project.create_log_stream(name="Production Logs") + log_stream = project.create_agent_stream(name="Production Logs") # List log streams for the project - log_streams = project.list_log_streams() + log_streams = project.list_agent_streams() # Access related resources via properties - for log_stream in project.logstreams: - print(log_stream.name) + for stream in project.agent_streams: + print(stream.name) for experiment in project.experiments: print(experiment.name) @@ -305,41 +304,11 @@ def create_agent_stream(self, name: str) -> "AgentStream": project = Project.get(name="My AI Project") stream = project.create_agent_stream(name="Production Traces") """ - from splunk_ao.agent_stream import AgentStream # lazy to avoid circular import - if self.id is None: raise ValueError("Project ID is not set. Cannot create agent stream for a local-only project.") - # cast: LogStream.create() returns LogStream, but self is AgentStream so - # the runtime type is correct. cast tells mypy to trust us. - return cast("AgentStream", AgentStream(name=name, project_id=self.id).create()) - - def create_log_stream(self, name: str) -> "AgentStream": - """ - Create a new agent stream for this project. - - .. deprecated:: - Use :meth:`create_agent_stream` instead. - - Args: - name (str): The name of the agent stream to create. - - Returns - ------- - AgentStream: The created agent stream. + return AgentStream(name=name, project_id=self.id).create() - Examples - -------- - project = Project.get(name="My AI Project") - log_stream = project.create_log_stream(name="Production Logs") - """ - import warnings - warnings.warn( - "Project.create_log_stream() is deprecated; use Project.create_agent_stream() instead.", - DeprecationWarning, - stacklevel=2, - ) - return self.create_agent_stream(name=name) def list_agent_streams( self, *, limit: Unset | int = 100, starting_token: Unset | int = 0 @@ -365,42 +334,11 @@ def list_agent_streams( for stream in streams: pass """ - from splunk_ao.agent_stream import AgentStream # lazy to avoid circular import - if self.id is None: raise ValueError("Project ID is not set. Cannot list agent streams for a local-only project.") - # cast: LogStream.list() is typed -> list[LogStream]; runtime values are - # AgentStream instances because cls is AgentStream. - return cast("builtins.list[AgentStream]", AgentStream.list(project_id=self.id, limit=limit, starting_token=starting_token)) + return AgentStream.list(project_id=self.id, limit=limit, starting_token=starting_token) - def list_log_streams( - self, *, limit: Unset | int = 100, starting_token: Unset | int = 0 - ) -> "builtins.list[AgentStream]": - """ - List agent streams for this project. - - .. deprecated:: - Use :meth:`list_agent_streams` instead. - - Returns a single page of results. Use `starting_token` (from - `next_starting_token` on a prior response) to fetch subsequent pages. - - Args: - limit (Union[Unset, int]): Maximum number of streams per page. Defaults to 100. - starting_token (Union[Unset, int]): Pagination token. Defaults to 0 (first page). - - Returns - ------- - List[AgentStream]: A page of agent streams belonging to this project. - """ - import warnings - warnings.warn( - "Project.list_log_streams() is deprecated; use Project.list_agent_streams() instead.", - DeprecationWarning, - stacklevel=2, - ) - return self.list_agent_streams(limit=limit, starting_token=starting_token) def list_experiments(self) -> builtins.list[Experiment]: """ @@ -488,25 +426,6 @@ def agent_streams(self) -> "builtins.list[AgentStream]": """ return self.list_agent_streams() - @property - def logstreams(self) -> "builtins.list[AgentStream]": - """ - Property to access agent streams for this project. - - .. deprecated:: - Use :attr:`agent_streams` instead. - - Returns - ------- - List[AgentStream]: A list of agent streams belonging to this project. - """ - import warnings - warnings.warn( - "Project.logstreams is deprecated; use Project.agent_streams instead.", - DeprecationWarning, - stacklevel=2, - ) - return self.list_agent_streams() @property def experiments(self) -> builtins.list[Experiment]: @@ -965,11 +884,8 @@ def save(self) -> Project: return self -# Import at end to avoid circular import (log_stream.py imports Project) +# Import at end to avoid circular import (agent_stream.py imports Project) +from splunk_ao.agent_stream import AgentStream # noqa: E402 from splunk_ao.dataset import Dataset # noqa: E402 from splunk_ao.experiment import Experiment # noqa: E402 -from splunk_ao.log_stream import LogStream # noqa: E402 from splunk_ao.prompt import Prompt # noqa: E402 - -# AgentStream is imported lazily inside methods to avoid a secondary circular import: -# agent_stream → log_stream → project → agent_stream diff --git a/src/splunk_ao/types.py b/src/splunk_ao/types.py index eec93cd6..ec2d610d 100644 --- a/src/splunk_ao/types.py +++ b/src/splunk_ao/types.py @@ -5,13 +5,13 @@ and other Galileo objects. """ -from splunk_ao.metric import Metric +from splunk_ao.evaluator import Evaluator from splunk_ao.schema.metrics import LocalMetricConfig, SplunkAOMetrics # Unified metric type that accepts all valid metric specifications MetricSpec = ( SplunkAOMetrics # Built-in scorer enum (e.g., SplunkAOMetrics.correctness) - | Metric # Custom or local metric object + | Evaluator # Custom or local evaluator object | LocalMetricConfig # Legacy local metric config | str # String name of built-in metric (e.g., "correctness") ) diff --git a/tests/test_agent_control_bridge.py b/tests/test_agent_control_bridge.py index ce17dfe0..bc6706c4 100644 --- a/tests/test_agent_control_bridge.py +++ b/tests/test_agent_control_bridge.py @@ -139,7 +139,7 @@ def _make_event(logger: SplunkAOLogger, **overrides: object) -> FakeControlExecu return FakeControlExecutionEvent(**payload) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_enable_agent_control_registers_provider_and_sink( @@ -171,7 +171,7 @@ def test_enable_agent_control_registers_provider_and_sink( assert fake_agent_control_modules["trace_context"].get_trace_context_from_provider() is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_logger_auto_registers_agent_control_bridge_when_available( @@ -191,7 +191,7 @@ def test_logger_auto_registers_agent_control_bridge_when_available( assert fake_agent_control_modules["agent_control"]._registered_sinks == [bridge._sink] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_logger_init_does_not_raise_when_agent_control_is_missing( @@ -212,7 +212,7 @@ def test_logger_init_does_not_raise_when_agent_control_is_missing( assert getattr(logger, "_agent_control_bridge", None) is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_agent_control_cleanup_restores_previous_provider_across_loggers( @@ -257,7 +257,7 @@ def external_provider(): } -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_agent_control_cleanup_does_not_clobber_provider_installed_while_active( @@ -290,7 +290,7 @@ def replacement_provider(): assert fake_agent_control_modules["agent_control"]._registered_sinks == [] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_idle_new_logger_does_not_mask_active_logger_context( @@ -322,7 +322,7 @@ def test_idle_new_logger_does_not_mask_active_logger_context( assert len(workflow_a.spans) == 1 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_agent_control_event_converts_to_control_span_in_batch_mode( @@ -371,7 +371,7 @@ def test_agent_control_event_converts_to_control_span_in_batch_mode( assert flushed_control_span.control_id == 7 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_agent_control_event_uses_empty_string_when_no_representative_input( @@ -396,7 +396,7 @@ def test_agent_control_event_uses_empty_string_when_no_representative_input( assert workflow.spans[0].input == "" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_agent_control_event_is_dropped_when_ids_are_not_valid_uuids( @@ -421,7 +421,7 @@ def test_agent_control_event_is_dropped_when_ids_are_not_valid_uuids( assert workflow.spans == [] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_agent_control_event_streams_immediately_in_distributed_mode( @@ -468,7 +468,7 @@ def test_agent_control_event_streams_immediately_in_distributed_mode( mock_traces_client_instance.ingest_spans.assert_called_with(request) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_add_control_span_uses_model_default_name( @@ -490,7 +490,7 @@ def test_add_control_span_uses_model_default_name( assert workflow.spans == [control_span] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_agent_control_event_is_dropped_when_context_does_not_match( diff --git a/tests/test_log_stream.py b/tests/test_agent_stream.py similarity index 89% rename from tests/test_log_stream.py rename to tests/test_agent_stream.py index 343ca3ea..59e8cab1 100644 --- a/tests/test_log_stream.py +++ b/tests/test_agent_stream.py @@ -4,7 +4,7 @@ import pytest from splunk_ao.exceptions import NotFoundError -from splunk_ao.log_stream import LogStream +from splunk_ao.agent_stream import AgentStream from splunk_ao.projects import ProjectNotFoundError, ProjectsAPIException from splunk_ao.resources.models import LLMExportFormat, LogRecordsSortClause, RootType from splunk_ao.resources.models.log_records_column_info import LogRecordsColumnInfo @@ -17,12 +17,12 @@ class TestLogStreamInitialization: - """Test suite for LogStream initialization.""" + """Test suite for AgentStream initialization.""" @pytest.mark.parametrize("project_kwarg", [{"project_id": "test-project-id"}, {"project_name": "Test Project"}]) def test_init_with_name_and_project(self, project_kwarg: dict, reset_configuration: None) -> None: """Test initializing a log stream with name and project creates a local-only instance.""" - log_stream = LogStream(name="Test Stream", **project_kwarg) + log_stream = AgentStream(name="Test Stream", **project_kwarg) assert log_stream.name == "Test Stream" assert log_stream.id is None @@ -38,13 +38,13 @@ def test_init_with_name_and_project(self, project_kwarg: dict, reset_configurati def test_init_without_name_raises_validation_error(self, reset_configuration: None) -> None: """Test initializing a log stream without a name raises ValidationError.""" with pytest.raises(ValidationError, match="'name' must be provided"): - LogStream(name="", project_id="test-project-id") + AgentStream(name="", project_id="test-project-id") def test_init_without_project_succeeds(self, reset_configuration: None) -> None: """Test initializing a log stream without project info succeeds (validated at create time).""" # Given: no project_id or project_name provided # When: creating a log stream - log_stream = LogStream(name="Test Stream") + log_stream = AgentStream(name="Test Stream") # Then: log stream is created with LOCAL_ONLY state, project info is None assert log_stream.project_id is None @@ -55,7 +55,7 @@ def test_init_with_both_project_id_and_name_succeeds(self, reset_configuration: """Test initializing a log stream with both project_id and project_name succeeds.""" # Given: both project_id and project_name provided # When: creating a log stream - log_stream = LogStream(name="Test Stream", project_id="test-id", project_name="Test Project") + log_stream = AgentStream(name="Test Stream", project_id="test-id", project_name="Test Project") # Then: log stream is created with both values stored assert log_stream.project_id == "test-id" @@ -63,9 +63,9 @@ def test_init_with_both_project_id_and_name_succeeds(self, reset_configuration: class TestLogStreamCreate: - """Test suite for LogStream.create() method.""" + """Test suite for AgentStream.create() method.""" - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_create_persists_log_stream_to_api_with_project_id( self, @@ -88,14 +88,14 @@ def test_create_persists_log_stream_to_api_with_project_id( mock_service.create.return_value = mock_logstream # When: creating log stream with project_id - log_stream = LogStream(name="Test Stream", project_id="test-project-id").create() + log_stream = AgentStream(name="Test Stream", project_id="test-project-id").create() # Then: log stream is created with resolved project_id - mock_service.create.assert_called_once_with(name="Test Stream", project_id="test-project-id", project_name=None) + mock_service.create.assert_called_once_with(name="Test Stream", project_id="test-project-id") assert log_stream.id == mock_logstream.id assert log_stream.is_synced() - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_create_persists_log_stream_to_api_with_project_name( self, @@ -118,17 +118,17 @@ def test_create_persists_log_stream_to_api_with_project_name( mock_service.create.return_value = mock_logstream # When: creating log stream with project_name - log_stream = LogStream(name="Test Stream", project_name="Test Project").create() + log_stream = AgentStream(name="Test Stream", project_name="Test Project").create() # Then: log stream is created with resolved project_id mock_service.create.assert_called_once_with( - name="Test Stream", project_id="resolved-project-id", project_name=None + name="Test Stream", project_id="resolved-project-id" ) assert log_stream.id == mock_logstream.id assert log_stream.is_synced() assert log_stream.project_name == "Test Project" - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_create_handles_api_failure( self, mock_projects_class: MagicMock, mock_logstreams_class: MagicMock, reset_configuration: None @@ -147,7 +147,7 @@ def test_create_handles_api_failure( mock_service.create.side_effect = Exception("API Error") # When/Then: create() raises error and sets FAILED_SYNC state - log_stream = LogStream(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream(name="Test Stream", project_id="test-project-id") with pytest.raises(Exception, match="API Error"): log_stream.create() @@ -165,7 +165,7 @@ def test_create_names_project_in_error_when_project_name_not_found( mock_projects_service.get_with_env_fallbacks.return_value = None # When: creating a log stream with a project_name that doesn't exist on the server - log_stream = LogStream(name="Test Stream", project_name="my-nonexistent-project") + log_stream = AgentStream(name="Test Stream", project_name="my-nonexistent-project") # Then: error names the project the user specified, not generic guidance with pytest.raises(NotFoundError, match=r'Project "my-nonexistent-project" not found'): @@ -182,7 +182,7 @@ def test_create_without_project_info_raises_error( mock_projects_service.get_with_env_fallbacks.return_value = None # Manually create instance to bypass __init__ validation - log_stream = LogStream._create_empty() + log_stream = AgentStream._create_empty() log_stream.name = "Test Stream" log_stream.project_id = None log_stream.project_name = None @@ -194,9 +194,9 @@ def test_create_without_project_info_raises_error( class TestLogStreamGet: - """Test suite for LogStream.get() class method.""" + """Test suite for AgentStream.get() class method.""" - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_get_returns_log_stream_with_project_id( self, @@ -219,7 +219,7 @@ def test_get_returns_log_stream_with_project_id( mock_service.get.return_value = mock_logstream # When: calling get with project_id - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") # Then: log stream is returned and project_name is set from resolved project assert log_stream is not None @@ -227,7 +227,7 @@ def test_get_returns_log_stream_with_project_id( assert log_stream.project_name == "Test Project" mock_service.get.assert_called_once_with(name="Test Stream", project_id="test-project-id") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_get_returns_log_stream_with_project_name( self, @@ -250,7 +250,7 @@ def test_get_returns_log_stream_with_project_name( mock_service.get.return_value = mock_logstream # When: calling get with project_name - log_stream = LogStream.get(name="Test Stream", project_name="Test Project") + log_stream = AgentStream.get(name="Test Stream", project_name="Test Project") # Then: log stream is returned using resolved project_id assert log_stream is not None @@ -258,7 +258,7 @@ def test_get_returns_log_stream_with_project_name( assert log_stream.project_name == "Test Project" mock_service.get.assert_called_once_with(name="Test Stream", project_id="resolved-project-id") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_get_returns_none_when_not_found( self, mock_projects_class: MagicMock, mock_logstreams_class: MagicMock, reset_configuration: None @@ -277,7 +277,7 @@ def test_get_returns_none_when_not_found( mock_service.get.return_value = None # When: calling get - log_stream = LogStream.get(name="Nonexistent Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Nonexistent Stream", project_id="test-project-id") # Then: None is returned assert log_stream is None @@ -294,7 +294,7 @@ def test_get_raises_error_without_project_info_and_no_env_fallback( # When/Then: Calling get raises NotFoundError with guidance to provide a project identifier with pytest.raises(NotFoundError, match="No project specified"): - LogStream.get(name="Test Stream") + AgentStream.get(name="Test Stream") @patch("splunk_ao.shared.project_resolver.Projects") def test_get_raises_not_found_when_project_id_unknown( @@ -308,7 +308,7 @@ def test_get_raises_not_found_when_project_id_unknown( # When/Then: calling get with an unknown project_id raises NotFoundError with the id in the message with pytest.raises(NotFoundError, match=r'Project with id "unknown-id" not found'): - LogStream.get(name="Test Stream", project_id="unknown-id") + AgentStream.get(name="Test Stream", project_id="unknown-id") @patch("splunk_ao.shared.project_resolver.Projects") def test_get_reraises_non_404_projects_api_exception( @@ -322,9 +322,9 @@ def test_get_reraises_non_404_projects_api_exception( # When/Then: non-404 errors propagate unchanged so callers receive the correct exception with pytest.raises(ProjectsAPIException): - LogStream.get(name="Test Stream", project_id="some-id") + AgentStream.get(name="Test Stream", project_id="some-id") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_get_uses_env_fallback_when_no_project_specified( self, @@ -349,7 +349,7 @@ def test_get_uses_env_fallback_when_no_project_specified( mock_service.get.return_value = mock_logstream # When: calling get without project params - log_stream = LogStream.get(name="Test Stream") + log_stream = AgentStream.get(name="Test Stream") # Then: project is resolved from env fallbacks mock_projects_service.get_with_env_fallbacks.assert_called_once() @@ -357,9 +357,9 @@ def test_get_uses_env_fallback_when_no_project_specified( class TestLogStreamList: - """Test suite for LogStream.list() class method.""" + """Test suite for AgentStream.list() class method.""" - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_list_returns_all_log_streams_with_project_id( self, mock_projects_class: MagicMock, mock_logstreams_class: MagicMock, reset_configuration: None @@ -391,16 +391,16 @@ def test_list_returns_all_log_streams_with_project_id( mock_service.list.return_value = mock_logstreams # When: calling list with project_id - log_streams = LogStream.list(project_id="test-project-id") + log_streams = AgentStream.list(project_id="test-project-id") # Then: all log streams are returned with project_name set assert len(log_streams) == 3 - assert all(isinstance(ls, LogStream) for ls in log_streams) + assert all(isinstance(ls, AgentStream) for ls in log_streams) assert all(ls.is_synced() for ls in log_streams) assert all(ls.project_name == "Test Project" for ls in log_streams) mock_service.list.assert_called_once_with(project_id="test-project-id", limit=100, starting_token=0) - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_list_returns_all_log_streams_with_project_name( self, mock_projects_class: MagicMock, mock_logstreams_class: MagicMock, reset_configuration: None @@ -428,7 +428,7 @@ def test_list_returns_all_log_streams_with_project_name( mock_service.list.return_value = [mock_ls] # When: calling list with project_name - log_streams = LogStream.list(project_name="Test Project") + log_streams = AgentStream.list(project_name="Test Project") # Then: log streams are returned using resolved project_id assert all(ls.project_name == "Test Project" for ls in log_streams) @@ -446,7 +446,7 @@ def test_list_raises_error_without_project_info_and_no_env_fallback( # When/Then: Calling list raises NotFoundError with guidance to provide a project identifier with pytest.raises(NotFoundError, match="No project specified"): - LogStream.list() + AgentStream.list() @patch("splunk_ao.shared.project_resolver.Projects") def test_list_raises_not_found_when_project_id_unknown( @@ -460,7 +460,7 @@ def test_list_raises_not_found_when_project_id_unknown( # When/Then: calling list with an unknown project_id raises NotFoundError with the id in the message with pytest.raises(NotFoundError, match=r'Project with id "unknown-id" not found'): - LogStream.list(project_id="unknown-id") + AgentStream.list(project_id="unknown-id") @patch("splunk_ao.shared.project_resolver.Projects") def test_list_reraises_non_404_projects_api_exception( @@ -474,9 +474,9 @@ def test_list_reraises_non_404_projects_api_exception( # When/Then: non-404 errors propagate unchanged so callers receive the correct exception with pytest.raises(ProjectsAPIException): - LogStream.list(project_id="some-id") + AgentStream.list(project_id="some-id") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_list_forwards_limit_to_service( self, mock_projects_class: MagicMock, mock_logstreams_class: MagicMock, reset_configuration: None @@ -495,12 +495,12 @@ def test_list_forwards_limit_to_service( mock_service.list.return_value = [] # When: calling list with a custom limit - LogStream.list(project_id="test-project-id", limit=3) + AgentStream.list(project_id="test-project-id", limit=3) # Then: limit is forwarded to the service call mock_service.list.assert_called_once_with(project_id="test-project-id", limit=3, starting_token=0) - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_list_forwards_starting_token_to_service( self, mock_projects_class: MagicMock, mock_logstreams_class: MagicMock, reset_configuration: None @@ -519,12 +519,12 @@ def test_list_forwards_starting_token_to_service( mock_service.list.return_value = [] # When: calling list with a custom starting_token - LogStream.list(project_id="test-project-id", starting_token=100) + AgentStream.list(project_id="test-project-id", starting_token=100) # Then: starting_token is forwarded to the service call mock_service.list.assert_called_once_with(project_id="test-project-id", limit=100, starting_token=100) - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") @patch("splunk_ao.shared.project_resolver.Projects") def test_list_uses_env_fallback_when_no_project_specified( self, @@ -548,7 +548,7 @@ def test_list_uses_env_fallback_when_no_project_specified( mock_service.list.return_value = [] # When: calling list without project params - LogStream.list() + AgentStream.list() # Then: project is resolved from env fallbacks mock_projects_service.get_with_env_fallbacks.assert_called_once() @@ -556,10 +556,10 @@ def test_list_uses_env_fallback_when_no_project_specified( class TestLogStreamRefresh: - """Test suite for LogStream.refresh() method.""" + """Test suite for AgentStream.refresh() method.""" @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") def test_refresh_updates_attributes_from_api( self, mock_logstreams_class: MagicMock, @@ -593,7 +593,7 @@ def test_refresh_updates_attributes_from_api( mock_service.get.side_effect = [initial_stream, updated_stream] - log_stream = LogStream.get(name="Old Name", project_id="test-project-id") + log_stream = AgentStream.get(name="Old Name", project_id="test-project-id") assert log_stream.name == "Old Name" log_stream.refresh() @@ -604,13 +604,13 @@ def test_refresh_updates_attributes_from_api( def test_refresh_raises_error_for_local_only(self, reset_configuration: None) -> None: """Test refresh() raises ValueError for local-only log stream.""" - log_stream = LogStream(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream(name="Test Stream", project_id="test-project-id") with pytest.raises(ValueError, match="Log stream ID is not set"): log_stream.refresh() @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") def test_refresh_raises_error_if_log_stream_no_longer_exists( self, mock_logstreams_class: MagicMock, @@ -625,20 +625,20 @@ def test_refresh_raises_error_if_log_stream_no_longer_exists( mock_logstreams_class.return_value = mock_service mock_service.get.side_effect = [mock_logstream, None] - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") with pytest.raises(ValueError, match="no longer exists"): log_stream.refresh() assert log_stream.sync_state == SyncState.FAILED_SYNC - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") def test_refresh_without_project_id_raises_error( self, mock_logstreams_class: MagicMock, reset_configuration: None, mock_logstream: MagicMock ) -> None: """Test refresh() raises ValueError when project_id is not set.""" # Manually create instance with id but no project_id - log_stream = LogStream._create_empty() + log_stream = AgentStream._create_empty() log_stream.id = str(uuid4()) log_stream.name = "Test Stream" log_stream.project_id = None @@ -649,7 +649,7 @@ def test_refresh_without_project_id_raises_error( class TestLogStreamQuery: - """Test suite for LogStream.query() and related methods.""" + """Test suite for AgentStream.query() and related methods.""" @pytest.mark.parametrize( "method_name,record_type,limit", @@ -663,8 +663,8 @@ class TestLogStreamQuery: ], ) @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.Search") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.Search") + @patch("splunk_ao.agent_stream.AgentStreams") def test_query_methods( self, mock_logstreams_class: MagicMock, @@ -688,7 +688,7 @@ def test_query_methods( mock_response = MagicMock() mock_search.query.return_value = mock_response - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") # Call the appropriate method if method_name == "query": @@ -712,18 +712,18 @@ def test_query_methods( def test_query_raises_error_for_local_only(self, reset_configuration: None) -> None: """Test query() raises ValueError for local-only log stream.""" - log_stream = LogStream(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream(name="Test Stream", project_id="test-project-id") with pytest.raises(ValueError, match="Log stream ID is not set"): log_stream.query(record_type=RecordType.SPAN) - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") def test_query_raises_error_without_project_id( self, mock_logstreams_class: MagicMock, reset_configuration: None, mock_logstream: MagicMock ) -> None: """Test query() raises ValueError when project_id is not set.""" # Manually create instance with id but no project_id - log_stream = LogStream._create_empty() + log_stream = AgentStream._create_empty() log_stream.id = str(uuid4()) log_stream.name = "Test Stream" log_stream.project_id = None @@ -734,11 +734,11 @@ def test_query_raises_error_without_project_id( class TestLogStreamExportRecords: - """Test suite for LogStream.export_records() method.""" + """Test suite for AgentStream.export_records() method.""" @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.ExportClient") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.ExportClient") + @patch("splunk_ao.agent_stream.AgentStreams") def test_export_records_with_default_params( self, mock_logstreams_class: MagicMock, @@ -759,7 +759,7 @@ def test_export_records_with_default_params( mock_iterator = iter([{"data": "test"}]) mock_export_client.records.return_value = mock_iterator - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") result = log_stream.export_records() # Verify ExportClient.records was called with correct parameters @@ -776,8 +776,8 @@ def test_export_records_with_default_params( assert result == mock_iterator @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.ExportClient") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.ExportClient") + @patch("splunk_ao.agent_stream.AgentStreams") def test_export_records_with_custom_params( self, mock_logstreams_class: MagicMock, @@ -798,7 +798,7 @@ def test_export_records_with_custom_params( mock_iterator = iter([{"data": "test"}]) mock_export_client.records.return_value = mock_iterator - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") custom_sort = LogRecordsSortClause(column_id="updated_at", ascending=True) log_stream.export_records( record_type=RecordType.SPAN, @@ -822,18 +822,18 @@ def test_export_records_with_custom_params( def test_export_records_raises_error_for_local_only(self, reset_configuration: None) -> None: """Test export_records() raises ValueError for local-only log stream.""" - log_stream = LogStream(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream(name="Test Stream", project_id="test-project-id") with pytest.raises(ValueError, match="Log stream ID is not set"): log_stream.export_records() - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.AgentStreams") def test_export_records_raises_error_without_project_id( self, mock_logstreams_class: MagicMock, reset_configuration: None ) -> None: """Test export_records() raises ValueError when project_id is not set.""" # Manually create instance with id but no project_id - log_stream = LogStream._create_empty() + log_stream = AgentStream._create_empty() log_stream.id = str(uuid4()) log_stream.name = "Test Stream" log_stream.project_id = None @@ -844,11 +844,11 @@ def test_export_records_raises_error_without_project_id( class TestLogStreamContext: - """Test suite for LogStream.context() method.""" + """Test suite for AgentStream.context() method.""" @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.splunk_ao_context") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.splunk_ao_context") + @patch("splunk_ao.agent_stream.AgentStreams") def test_context_returns_splunk_ao_context( self, mock_logstreams_class: MagicMock, @@ -868,12 +868,12 @@ def test_context_returns_splunk_ao_context( mock_splunk_ao_context.return_value = mock_context # Mock the project property - with patch("splunk_ao.log_stream.Project") as mock_project_class: + with patch("splunk_ao.agent_stream.Project") as mock_project_class: mock_project = MagicMock() mock_project.name = "Test Project" mock_project_class.get.return_value = mock_project - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") result = log_stream.context() mock_splunk_ao_context.assert_called_once_with(project="Test Project", log_stream="Test Stream") @@ -881,11 +881,11 @@ def test_context_returns_splunk_ao_context( class TestLogStreamProject: - """Test suite for LogStream.project property.""" + """Test suite for AgentStream.project property.""" @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.Project") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.Project") + @patch("splunk_ao.agent_stream.AgentStreams") def test_project_property_returns_project( self, mock_logstreams_class: MagicMock, @@ -906,7 +906,7 @@ def test_project_property_returns_project( returned_project.name = "Test Project" mock_project_class.get.return_value = returned_project - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") project = log_stream.project mock_project_class.get.assert_called_once_with(id=mock_logstream.project_id) @@ -914,7 +914,7 @@ def test_project_property_returns_project( class TestLogStreamColumns: - """Test suite for LogStream column properties.""" + """Test suite for AgentStream column properties.""" @pytest.mark.parametrize( "property_name,api_func_name,error_msg", @@ -937,8 +937,8 @@ class TestLogStreamColumns: ], ) @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.SplunkAOConfig") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.SplunkAOConfig") + @patch("splunk_ao.agent_stream.AgentStreams") def test_column_properties_return_column_collection( self, mock_logstreams_class: MagicMock, @@ -953,7 +953,7 @@ def test_column_properties_return_column_collection( ) -> None: mock_projects_class.return_value.get_with_env_fallbacks.return_value = mock_project """Test column properties return ColumnCollection with proper API calls.""" - # Setup LogStreams mock + # Setup AgentStreams mock mock_logstream_service = MagicMock() mock_logstreams_class.return_value = mock_logstream_service mock_logstream_service.get.return_value = mock_logstream @@ -972,10 +972,10 @@ def test_column_properties_return_column_collection( mock_response = MagicMock() mock_response.columns = [mock_column_1, mock_column_2] - with patch(f"splunk_ao.log_stream.{api_func_name}") as mock_api_func: + with patch(f"splunk_ao.agent_stream.{api_func_name}") as mock_api_func: mock_api_func.sync.return_value = mock_response - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") columns = getattr(log_stream, property_name) # Verify API function was called correctly @@ -993,7 +993,7 @@ def test_column_properties_return_column_collection( @pytest.mark.parametrize("property_name", ["span_columns", "session_columns", "trace_columns"]) def test_column_properties_raise_error_for_local_only(self, property_name: str, reset_configuration: None) -> None: """Test column properties raise ValueError for local-only log streams.""" - log_stream = LogStream(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream(name="Test Stream", project_id="test-project-id") with pytest.raises(ValueError, match="Log stream ID is not set"): getattr(log_stream, property_name) @@ -1007,8 +1007,8 @@ def test_column_properties_raise_error_for_local_only(self, property_name: str, ], ) @patch("splunk_ao.shared.project_resolver.Projects") - @patch("splunk_ao.log_stream.SplunkAOConfig") - @patch("splunk_ao.log_stream.LogStreams") + @patch("splunk_ao.agent_stream.SplunkAOConfig") + @patch("splunk_ao.agent_stream.AgentStreams") def test_column_properties_raise_error_on_empty_response( self, mock_logstreams_class: MagicMock, @@ -1030,10 +1030,10 @@ def test_column_properties_raise_error_on_empty_response( mock_config = MagicMock() mock_config_class.get.return_value = mock_config - with patch(f"splunk_ao.log_stream.{api_func_name}") as mock_api_func: + with patch(f"splunk_ao.agent_stream.{api_func_name}") as mock_api_func: mock_api_func.sync.return_value = None - log_stream = LogStream.get(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream.get(name="Test Stream", project_id="test-project-id") with pytest.raises(ValueError, match="Unable to retrieve"): getattr(log_stream, property_name) @@ -1057,18 +1057,18 @@ def test_column_info_with_control_step_type_does_not_raise(self, reset_configura class TestLogStreamMethods: - """Test suite for other LogStream methods.""" + """Test suite for other AgentStream methods.""" def test_str_representation(self, reset_configuration: None) -> None: """Test __str__ returns expected format.""" - log_stream = LogStream(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream(name="Test Stream", project_id="test-project-id") log_stream.id = "test-id-123" - assert str(log_stream) == "LogStream(name='Test Stream', id='test-id-123', project_id='test-project-id')" + assert str(log_stream) == "AgentStream(name='Test Stream', id='test-id-123', project_id='test-project-id')" def test_repr_representation(self, reset_configuration: None) -> None: """Test __repr__ returns expected format with created_at.""" - log_stream = LogStream(name="Test Stream", project_id="test-project-id") + log_stream = AgentStream(name="Test Stream", project_id="test-project-id") log_stream.id = "test-id-123" log_stream.created_at = "2024-01-01 12:00:00" @@ -1094,7 +1094,7 @@ def test_create_raises_resource_not_found_error_subclass( mock_projects_class.return_value = mock_projects_service mock_projects_service.get_with_env_fallbacks.return_value = None - log_stream = LogStream(name="Test Stream", project_name="missing-proj") + log_stream = AgentStream(name="Test Stream", project_name="missing-proj") # When/Then: the raised exception is BOTH a NotFoundError AND a ResourceNotFoundError with pytest.raises(NotFoundError) as exc_info: @@ -1117,7 +1117,7 @@ def test_create_skips_api_when_no_identifier_anywhere( monkeypatch.delenv("SPLUNK_AO_PROJECT", raising=False) monkeypatch.delenv("SPLUNK_AO_PROJECT_ID", raising=False) - log_stream = LogStream._create_empty() + log_stream = AgentStream._create_empty() log_stream.name = "Test Stream" log_stream.project_id = None log_stream.project_name = None @@ -1143,7 +1143,7 @@ def test_resolver_does_not_swallow_unrelated_value_error( mock_projects_class.return_value = mock_projects_service mock_projects_service.get_with_env_fallbacks.side_effect = ValueError("HTTP client blew up") - log_stream = LogStream._create_empty() + log_stream = AgentStream._create_empty() log_stream.name = "Test Stream" log_stream.project_id = "explicit-id" log_stream.project_name = None diff --git a/tests/test_log_streams_metrics.py b/tests/test_agent_streams_evaluators.py similarity index 82% rename from tests/test_log_streams_metrics.py rename to tests/test_agent_streams_evaluators.py index 1beac71d..e425d23a 100644 --- a/tests/test_log_streams_metrics.py +++ b/tests/test_agent_streams_evaluators.py @@ -4,7 +4,7 @@ import pytest -from splunk_ao.log_streams import LogStream, LogStreams, enable_metrics +from splunk_ao.agent_streams import AgentStream, AgentStreams, enable_evaluators from splunk_ao.projects import Project from splunk_ao.resources.models import ProjectCreateResponse, ScorerResponse, ScorerTypes from splunk_ao.resources.models.log_stream_response import LogStreamResponse @@ -40,7 +40,7 @@ def mock_project(): @pytest.fixture def mock_log_stream(): """Mock log stream object.""" - return LogStream( + return AgentStream( log_stream=LogStreamResponse( id="log-stream-123", name="Test Log Stream", @@ -154,12 +154,12 @@ def custom_scorer(trace_or_span) -> float: mock_scorer_settings_class.return_value.create.assert_called_once() def test_log_stream_enable_metrics_instance_method(self, mock_log_stream) -> None: - """Test LogStream instance enable_metrics method.""" - with patch("splunk_ao.log_streams.create_metric_configs") as mock_create_configs: + """Test AgentStream instance enable_evaluators method.""" + with patch("splunk_ao.agent_streams.create_metric_configs") as mock_create_configs: mock_create_configs.return_value = ([], []) # Test instance method - local_metrics = mock_log_stream.enable_metrics(["correctness"]) + local_metrics = mock_log_stream.enable_evaluators(["correctness"]) # Verify create_metric_configs was called with correct parameters mock_create_configs.assert_called_once_with(mock_log_stream.project_id, mock_log_stream.id, ["correctness"]) @@ -168,19 +168,19 @@ def test_log_stream_enable_metrics_instance_method(self, mock_log_stream) -> Non assert local_metrics == [] def test_log_stream_enable_metrics_missing_ids(self) -> None: - """Test LogStream enable_metrics raises error when IDs are missing.""" - log_stream = LogStream() # Empty log stream without IDs + """Test AgentStream enable_evaluators raises error when IDs are missing.""" + log_stream = AgentStream() # Empty log stream without IDs with pytest.raises(ValueError, match="Log stream must have id and project_id to enable metrics"): - log_stream.enable_metrics(["correctness"]) + log_stream.enable_evaluators(["correctness"]) - @patch("splunk_ao.log_streams.Projects") - @patch.object(LogStreams, "get") - @patch("splunk_ao.log_streams.create_metric_configs") + @patch("splunk_ao.agent_streams.Projects") + @patch.object(AgentStreams, "get") + @patch("splunk_ao.agent_streams.create_metric_configs") def test_logstreams_enable_metrics_with_explicit_params( self, mock_create_configs, mock_get, mock_projects_class, mock_project, mock_log_stream ) -> None: - """Test LogStreams.enable_metrics with explicit parameters.""" + """Test AgentStreams.enable_evaluators with explicit parameters.""" # Setup mocks mock_projects_instance = mock_projects_class.return_value mock_projects_instance.get_with_env_fallbacks.return_value = mock_project @@ -188,9 +188,9 @@ def test_logstreams_enable_metrics_with_explicit_params( mock_create_configs.return_value = ([], []) # Test with explicit parameters - log_streams = LogStreams() - local_metrics = log_streams.enable_metrics( - log_stream_name="Test Log Stream", project_name="Test Project", metrics=["correctness"] + log_streams = AgentStreams() + local_metrics = log_streams.enable_evaluators( + agent_stream_name="Test Log Stream", project_name="Test Project", metrics=["correctness"] ) # Verify project lookup @@ -205,22 +205,22 @@ def test_logstreams_enable_metrics_with_explicit_params( # Verify return value is just local metrics assert local_metrics == [] - @patch.object(LogStreams, "get") - @patch("splunk_ao.log_streams.create_metric_configs") + @patch.object(AgentStreams, "get") + @patch("splunk_ao.agent_streams.create_metric_configs") @patch("splunk_ao.projects.Projects.get_with_env_fallbacks") def test_logstreams_enable_metrics_gets_project_correctly( self, mock_get_with_env_fallbacks, mock_create_configs, mock_get, mock_log_stream, mock_project ) -> None: - """Test LogStreams.enable_metrics with explicit parameters.""" + """Test AgentStreams.enable_evaluators with explicit parameters.""" # Setup mocks mock_get.return_value = mock_log_stream mock_create_configs.return_value = ([], []) mock_get_with_env_fallbacks.return_value = mock_project # Test with explicit parameters - log_streams = LogStreams() - local_metrics = log_streams.enable_metrics( - log_stream_name="Test Log Stream", project_name="Test Project", metrics=["correctness"] + log_streams = AgentStreams() + local_metrics = log_streams.enable_evaluators( + agent_stream_name="Test Log Stream", project_name="Test Project", metrics=["correctness"] ) # Verify log stream lookup @@ -232,13 +232,13 @@ def test_logstreams_enable_metrics_gets_project_correctly( # Verify return value is just local metrics assert local_metrics == [] - @patch("splunk_ao.log_streams.Projects") - @patch.object(LogStreams, "get") - @patch("splunk_ao.log_streams.create_metric_configs") + @patch("splunk_ao.agent_streams.Projects") + @patch.object(AgentStreams, "get") + @patch("splunk_ao.agent_streams.create_metric_configs") def test_logstreams_enable_metrics_with_env_vars( self, mock_create_configs, mock_get, mock_projects_class, mock_project, mock_log_stream ) -> None: - """Test LogStreams.enable_metrics with environment variables.""" + """Test AgentStreams.enable_evaluators with environment variables.""" # Set environment variables os.environ["SPLUNK_AO_PROJECT"] = "Test Project" os.environ["SPLUNK_AO_LOG_STREAM"] = "Test Log Stream" @@ -250,8 +250,8 @@ def test_logstreams_enable_metrics_with_env_vars( mock_create_configs.return_value = ([], []) # Test with environment variables - log_streams = LogStreams() - local_metrics = log_streams.enable_metrics(metrics=["correctness"]) + log_streams = AgentStreams() + local_metrics = log_streams.enable_evaluators(metrics=["correctness"]) # Verify project lookup used env var mock_projects_instance.get_with_env_fallbacks.assert_called_once_with(name="Test Project") @@ -262,79 +262,79 @@ def test_logstreams_enable_metrics_with_env_vars( # Verify return value is just local metrics assert local_metrics == [] - @patch("splunk_ao.log_streams.Projects") + @patch("splunk_ao.agent_streams.Projects") def test_logstreams_enable_metrics_project_not_found(self, mock_projects_class) -> None: - """Test LogStreams.enable_metrics raises ValueError when project not found.""" + """Test AgentStreams.enable_evaluators raises ValueError when project not found.""" # Setup mock to return None mock_projects_instance = mock_projects_class.return_value mock_projects_instance.get_with_env_fallbacks.return_value = None # Test with non-existent project - expect ValueError - log_streams = LogStreams() + log_streams = AgentStreams() with pytest.raises(ValueError) as exc_info: - log_streams.enable_metrics( - project_name="Nonexistent Project", log_stream_name="Test Log Stream", metrics=["correctness"] + log_streams.enable_evaluators( + project_name="Nonexistent Project", agent_stream_name="Test Log Stream", metrics=["correctness"] ) assert "Project 'Nonexistent Project' not found" in str(exc_info.value) - @patch("splunk_ao.log_streams.Projects") - @patch.object(LogStreams, "get") + @patch("splunk_ao.agent_streams.Projects") + @patch.object(AgentStreams, "get") def test_logstreams_enable_metrics_logstream_not_found(self, mock_get, mock_projects_class, mock_project) -> None: - """Test LogStreams.enable_metrics raises ValueError when log stream not found.""" + """Test AgentStreams.enable_evaluators raises ValueError when log stream not found.""" # Setup mocks mock_projects_instance = mock_projects_class.return_value mock_projects_instance.get_with_env_fallbacks.return_value = mock_project mock_get.return_value = None # Log stream not found # Test with non-existent log stream - expect ValueError - log_streams = LogStreams() + log_streams = AgentStreams() with pytest.raises(ValueError) as exc_info: - log_streams.enable_metrics( - project_name="Test Project", log_stream_name="Nonexistent Stream", metrics=["correctness"] + log_streams.enable_evaluators( + project_name="Test Project", agent_stream_name="Nonexistent Stream", metrics=["correctness"] ) assert "Log stream 'Nonexistent Stream' not found" in str(exc_info.value) - @patch.object(LogStreams, "enable_metrics") + @patch.object(AgentStreams, "enable_evaluators") def test_enable_metrics_convenience_function_explicit(self, mock_enable_metrics) -> None: - """Test enable_metrics convenience function with explicit parameters.""" + """Test enable_evaluators convenience function with explicit parameters.""" mock_enable_metrics.return_value = [] # Test convenience function with explicit parameters - local_metrics = enable_metrics( - log_stream_name="Test Stream", project_name="Test Project", metrics=["correctness"] + local_metrics = enable_evaluators( + agent_stream_name="Test Stream", project_name="Test Project", metrics=["correctness"] ) # Verify it calls the instance method mock_enable_metrics.assert_called_once_with( - log_stream_name="Test Stream", project_name="Test Project", metrics=["correctness"] + agent_stream_name="Test Stream", project_name="Test Project", metrics=["correctness"] ) # Verify return value is just local metrics assert local_metrics == [] - @patch.object(LogStreams, "enable_metrics") + @patch.object(AgentStreams, "enable_evaluators") def test_enable_metrics_convenience_function_env_only(self, mock_enable_metrics) -> None: - """Test enable_metrics convenience function with environment variables only.""" + """Test enable_evaluators convenience function with environment variables only.""" mock_enable_metrics.return_value = [] # Test environment-only function (no explicit parameters) - local_metrics = enable_metrics(metrics=["correctness"]) + local_metrics = enable_evaluators(metrics=["correctness"]) # Verify it calls the instance method with no explicit params - mock_enable_metrics.assert_called_once_with(log_stream_name=None, project_name=None, metrics=["correctness"]) + mock_enable_metrics.assert_called_once_with(agent_stream_name=None, project_name=None, metrics=["correctness"]) # Verify return value is just local metrics assert local_metrics == [] - @patch("splunk_ao.log_streams.Projects") - @patch.object(LogStreams, "get") - @patch("splunk_ao.log_streams.create_metric_configs") + @patch("splunk_ao.agent_streams.Projects") + @patch.object(AgentStreams, "get") + @patch("splunk_ao.agent_streams.create_metric_configs") def test_enable_metrics_with_env_vars_integration( self, mock_create_configs, mock_get, mock_projects_class, mock_project, mock_log_stream ) -> None: - """Test enable_metrics function with environment variables (integration test).""" + """Test enable_evaluators function with environment variables (integration test).""" # Set environment variables os.environ["SPLUNK_AO_PROJECT"] = "Integration Project" os.environ["SPLUNK_AO_LOG_STREAM"] = "Integration Stream" @@ -346,7 +346,7 @@ def test_enable_metrics_with_env_vars_integration( mock_create_configs.return_value = ([], []) # Test the full integration - local_metrics = enable_metrics(metrics=["correctness", "completeness"]) + local_metrics = enable_evaluators(metrics=["correctness", "completeness"]) # Verify the entire chain was called correctly mock_projects_instance.get_with_env_fallbacks.assert_called_once_with(name="Integration Project") @@ -359,15 +359,15 @@ def test_enable_metrics_with_env_vars_integration( assert local_metrics == [] def test_enable_metrics_missing_env_vars(self) -> None: - """Test enable_metrics raises ValueError when environment variables are missing.""" + """Test enable_evaluators raises ValueError when environment variables are missing.""" # Don't set any environment variables - with patch("splunk_ao.log_streams.Projects") as mock_projects_class: + with patch("splunk_ao.agent_streams.Projects") as mock_projects_class: mock_projects_instance = mock_projects_class.return_value mock_projects_instance.get_with_env_fallbacks.return_value = None # Expect ValueError since project is not found with pytest.raises(ValueError) as exc_info: - enable_metrics(metrics=["correctness"]) + enable_evaluators(metrics=["correctness"]) assert "Project" in str(exc_info.value) assert "not found" in str(exc_info.value) diff --git a/tests/test_log_streams_pagination.py b/tests/test_agent_streams_pagination.py similarity index 71% rename from tests/test_log_streams_pagination.py rename to tests/test_agent_streams_pagination.py index 866fc7d7..73e32f99 100644 --- a/tests/test_log_streams_pagination.py +++ b/tests/test_agent_streams_pagination.py @@ -1,4 +1,4 @@ -"""Tests for pagination behavior of the LogStreams service. +"""Tests for pagination behavior of the AgentStreams service. Covers: - `_list_all` paginates across all pages until the API signals no more pages. @@ -12,7 +12,7 @@ import pytest -from splunk_ao.log_streams import LogStreams +from splunk_ao.agent_streams import AgentStreams from splunk_ao.resources.models.http_validation_error import HTTPValidationError from splunk_ao.resources.models.list_log_stream_response import ListLogStreamResponse from splunk_ao.resources.models.log_stream_response import LogStreamResponse @@ -38,10 +38,10 @@ def _make_response(*, names: list[str], next_token, paginated: bool) -> ListLogS class TestListAllPagination: - """Tests for LogStreams._list_all (internal helper).""" + """Tests for AgentStreams._list_all (internal helper).""" - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_all_paginates_across_multiple_pages( self, mock_config_class: MagicMock, mock_endpoint: MagicMock ) -> None: @@ -51,7 +51,7 @@ def test_list_all_paginates_across_multiple_pages( mock_endpoint.sync.side_effect = [page_1, page_2] # When: _list_all is called - all_streams = LogStreams()._list_all(project_id="proj-1") + all_streams = AgentStreams()._list_all(project_id="proj-1") # Then: every page is fetched and concatenated assert len(all_streams) == 8 @@ -60,22 +60,22 @@ def test_list_all_paginates_across_multiple_pages( # Second call passes the token from the first response assert mock_endpoint.sync.call_args_list[1].kwargs["starting_token"] == 5 - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_all_stops_when_paginated_false(self, mock_config_class: MagicMock, mock_endpoint: MagicMock) -> None: # Given: a single page with paginated=False page = _make_response(names=["only-stream"], next_token=42, paginated=False) mock_endpoint.sync.return_value = page # When: _list_all is called - all_streams = LogStreams()._list_all(project_id="proj-1") + all_streams = AgentStreams()._list_all(project_id="proj-1") # Then: only one fetch happens regardless of next_starting_token assert len(all_streams) == 1 assert mock_endpoint.sync.call_count == 1 - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_all_stops_when_next_token_is_unset( self, mock_config_class: MagicMock, mock_endpoint: MagicMock ) -> None: @@ -84,28 +84,28 @@ def test_list_all_stops_when_next_token_is_unset( mock_endpoint.sync.return_value = page # When: _list_all is called - all_streams = LogStreams()._list_all(project_id="proj-1") + all_streams = AgentStreams()._list_all(project_id="proj-1") # Then: pagination stops on UNSET token assert len(all_streams) == 2 assert mock_endpoint.sync.call_count == 1 - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_all_uses_larger_page_size(self, mock_config_class: MagicMock, mock_endpoint: MagicMock) -> None: # Given: a single page mock_endpoint.sync.return_value = _make_response(names=["a"], next_token=None, paginated=True) # When: _list_all is called - LogStreams()._list_all(project_id="proj-1") + AgentStreams()._list_all(project_id="proj-1") # Then: it uses the larger page size (500) to reduce round trips kwargs = mock_endpoint.sync.call_args.kwargs - assert kwargs["limit"] == LogStreams._LIST_ALL_PAGE_SIZE + assert kwargs["limit"] == AgentStreams._LIST_ALL_PAGE_SIZE assert kwargs["limit"] == 500 - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_all_raises_on_http_validation_error( self, mock_config_class: MagicMock, mock_endpoint: MagicMock ) -> None: @@ -115,20 +115,20 @@ def test_list_all_raises_on_http_validation_error( # When/Then: _list_all raises instead of silently returning the partial accumulator with pytest.raises(ValueError, match="Failed to list log streams"): - LogStreams()._list_all(project_id="proj-1") + AgentStreams()._list_all(project_id="proj-1") - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_all_raises_on_none_response(self, mock_config_class: MagicMock, mock_endpoint: MagicMock) -> None: # Given: the endpoint returns None (unexpected protocol error) mock_endpoint.sync.return_value = None # When/Then: _list_all raises ValueError instead of returning an empty list with pytest.raises(ValueError, match="Unexpected empty response"): - LogStreams()._list_all(project_id="proj-1") + AgentStreams()._list_all(project_id="proj-1") - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_all_breaks_on_non_advancing_token( self, mock_config_class: MagicMock, mock_endpoint: MagicMock ) -> None: @@ -137,14 +137,14 @@ def test_list_all_breaks_on_non_advancing_token( mock_endpoint.sync.return_value = same_token_page # When: _list_all is called - all_streams = LogStreams()._list_all(project_id="proj-1") + all_streams = AgentStreams()._list_all(project_id="proj-1") # Then: the loop terminates after one iteration thanks to the progress guard assert len(all_streams) == 1 assert mock_endpoint.sync.call_count == 1 - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_all_breaks_on_repeated_seen_token( self, mock_config_class: MagicMock, mock_endpoint: MagicMock ) -> None: @@ -157,7 +157,7 @@ def test_list_all_breaks_on_repeated_seen_token( mock_endpoint.sync.side_effect = [page_1, page_2, page_3] # When: _list_all is called - all_streams = LogStreams()._list_all(project_id="proj-1") + all_streams = AgentStreams()._list_all(project_id="proj-1") # Then: the loop terminates on the third page when the cycle is detected assert len(all_streams) == 3 @@ -165,10 +165,10 @@ def test_list_all_breaks_on_repeated_seen_token( class TestGetByNamePaginates: - """Tests for LogStreams.get(name=...) finding matches across pages.""" + """Tests for AgentStreams.get(name=...) finding matches across pages.""" - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_get_by_name_finds_match_on_second_page( self, mock_config_class: MagicMock, mock_endpoint: MagicMock ) -> None: @@ -178,15 +178,15 @@ def test_get_by_name_finds_match_on_second_page( mock_endpoint.sync.side_effect = [page_1, page_2] # When: looking up by name - result = LogStreams().get(name="target-stream", project_id="proj-1") + result = AgentStreams().get(name="target-stream", project_id="proj-1") # Then: the match on page 2 is found (would have returned None before this fix) assert result is not None assert result.name == "target-stream" assert mock_endpoint.sync.call_count == 2 - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_get_by_name_returns_none_when_missing( self, mock_config_class: MagicMock, mock_endpoint: MagicMock ) -> None: @@ -195,23 +195,23 @@ def test_get_by_name_returns_none_when_missing( mock_endpoint.sync.return_value = page # When: looking up a missing name - result = LogStreams().get(name="nonexistent", project_id="proj-1") + result = AgentStreams().get(name="nonexistent", project_id="proj-1") # Then: returns None assert result is None class TestListForwardsStartingToken: - """Tests that LogStreams.list forwards starting_token to the paginated endpoint.""" + """Tests that AgentStreams.list forwards starting_token to the paginated endpoint.""" - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_forwards_starting_token(self, mock_config_class: MagicMock, mock_endpoint: MagicMock) -> None: # Given: a single page response mock_endpoint.sync.return_value = _make_response(names=["s1"], next_token=None, paginated=True) # When: list is called with a custom starting_token - LogStreams().list(project_id="proj-1", starting_token=200, limit=50) + AgentStreams().list(project_id="proj-1", starting_token=200, limit=50) # Then: starting_token and limit are passed to the underlying endpoint kwargs = mock_endpoint.sync.call_args.kwargs @@ -219,14 +219,14 @@ def test_list_forwards_starting_token(self, mock_config_class: MagicMock, mock_e assert kwargs["limit"] == 50 assert kwargs["project_id"] == "proj-1" - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_default_starting_token_is_zero(self, mock_config_class: MagicMock, mock_endpoint: MagicMock) -> None: # Given: a single page response mock_endpoint.sync.return_value = _make_response(names=[], next_token=None, paginated=True) # When: list is called without starting_token - LogStreams().list(project_id="proj-1") + AgentStreams().list(project_id="proj-1") # Then: starting_token defaults to 0 kwargs = mock_endpoint.sync.call_args.kwargs @@ -235,38 +235,38 @@ def test_list_default_starting_token_is_zero(self, mock_config_class: MagicMock, class TestListValidatesArguments: - """Tests that LogStreams.list rejects invalid (project_id, project_name) combinations.""" + """Tests that AgentStreams.list rejects invalid (project_id, project_name) combinations.""" def test_list_raises_when_both_project_id_and_name_provided(self) -> None: # When/Then: passing both project_id and project_name is rejected (matches get/create XOR contract) with pytest.raises(ValueError, match="Exactly one of 'project_id' or 'project_name'"): - LogStreams().list(project_id="proj-1", project_name="My Project") + AgentStreams().list(project_id="proj-1", project_name="My Project") def test_list_raises_when_neither_project_id_nor_name_provided(self) -> None: # When/Then: passing neither is rejected with pytest.raises(ValueError, match="Exactly one of 'project_id' or 'project_name'"): - LogStreams().list() + AgentStreams().list() class TestListPropagatesErrors: - """Tests that LogStreams.list raises instead of silently returning [] on server errors.""" + """Tests that AgentStreams.list raises instead of silently returning [] on server errors.""" - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_raises_on_http_validation_error(self, mock_config_class: MagicMock, mock_endpoint: MagicMock) -> None: # Given: the endpoint returns an HTTPValidationError (e.g. bad starting_token type) mock_endpoint.sync.return_value = HTTPValidationError() # When/Then: list raises ValueError instead of masking the error as an empty page with pytest.raises(ValueError, match="Failed to list log streams"): - LogStreams().list(project_id="proj-1") + AgentStreams().list(project_id="proj-1") - @patch("splunk_ao.log_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") - @patch("splunk_ao.log_streams.SplunkAOConfig") + @patch("splunk_ao.agent_streams.list_log_streams_paginated_projects_project_id_log_streams_paginated_get") + @patch("splunk_ao.agent_streams.SplunkAOConfig") def test_list_raises_on_none_response(self, mock_config_class: MagicMock, mock_endpoint: MagicMock) -> None: # Given: the endpoint returns None (unexpected protocol error) mock_endpoint.sync.return_value = None # When/Then: list raises ValueError with pytest.raises(ValueError, match="Unexpected empty response"): - LogStreams().list(project_id="proj-1") + AgentStreams().list(project_id="proj-1") diff --git a/tests/test_async_base_handler.py b/tests/test_async_base_handler.py index f0d79e72..ed07dec4 100644 --- a/tests/test_async_base_handler.py +++ b/tests/test_async_base_handler.py @@ -11,7 +11,7 @@ class TestSplunkAOAsyncBaseHandlerCallback: @pytest.fixture - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def splunk_ao_logger(self, mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock): diff --git a/tests/test_backward_compat_future.py b/tests/test_backward_compat_future.py index cf664fc5..51bc298d 100644 --- a/tests/test_backward_compat_future.py +++ b/tests/test_backward_compat_future.py @@ -127,18 +127,18 @@ def test_provider_classes_are_same(): def test_metric_classes_are_same(): - from splunk_ao.__future__.metric import CodeMetric as FutureCodeMetric - from splunk_ao.__future__.metric import LlmMetric as FutureLlmMetric - from splunk_ao.__future__.metric import LocalMetric as FutureLocalMetric - from splunk_ao.__future__.metric import Metric as FutureMetric - from splunk_ao.__future__.metric import SplunkAOMetric as FutureSplunkAOMetric - from splunk_ao.metric import CodeMetric, LlmMetric, LocalMetric, Metric, SplunkAOMetric + from splunk_ao.__future__ import CodeEvaluator as FutureCodeMetric + from splunk_ao.__future__ import LlmEvaluator as FutureLlmMetric + from splunk_ao.__future__ import LocalEvaluator as FutureLocalMetric + from splunk_ao.__future__ import Evaluator as FutureMetric + from splunk_ao.__future__ import SplunkAOEvaluator as FutureSplunkAOMetric + from splunk_ao.evaluator import CodeEvaluator, LlmEvaluator, LocalEvaluator, Evaluator, SplunkAOEvaluator - assert FutureMetric is Metric - assert FutureCodeMetric is CodeMetric - assert FutureSplunkAOMetric is SplunkAOMetric - assert FutureLlmMetric is LlmMetric - assert FutureLocalMetric is LocalMetric + assert FutureMetric is Evaluator + assert FutureCodeMetric is CodeEvaluator + assert FutureSplunkAOMetric is SplunkAOEvaluator + assert FutureLlmMetric is LlmEvaluator + assert FutureLocalMetric is LocalEvaluator def test_experiment_is_same_class(): @@ -149,8 +149,8 @@ def test_experiment_is_same_class(): def test_log_stream_is_same_class(): - from splunk_ao.__future__.log_stream import LogStream as FutureLogStream - from splunk_ao.log_stream import LogStream as RootLogStream + from splunk_ao.__future__ import AgentStream as FutureLogStream + from splunk_ao.agent_stream import AgentStream as RootLogStream assert FutureLogStream is RootLogStream @@ -217,10 +217,10 @@ def test_provider_generic_and_unconfigured_are_same(): def test_metric_builtin_metrics_is_same(): - from splunk_ao.__future__.metric import BuiltInMetrics as FutureBuiltIn - from splunk_ao.metric import BuiltInMetrics + from splunk_ao.__future__ import BuiltInEvaluators as FutureBuiltIn + from splunk_ao.evaluator import BuiltInEvaluators - assert FutureBuiltIn is BuiltInMetrics + assert FutureBuiltIn is BuiltInEvaluators def test_prompt_private_symbols_are_same(): @@ -284,33 +284,33 @@ def test_root_init_has_new_exports(): AnthropicProvider, AzureProvider, BedrockProvider, - CodeMetric, + CodeEvaluator, Configuration, Dataset, Experiment, Integration, - LlmMetric, - LocalMetric, - LogStream, - Metric, + LlmEvaluator, + LocalEvaluator, + AgentStream, + Evaluator, MetricSpec, Model, OpenAIProvider, Prompt, Provider, - SplunkAOMetric, + SplunkAOEvaluator, ) assert Configuration is not None assert Dataset is not None assert Experiment is not None assert Integration is not None - assert LogStream is not None - assert Metric is not None - assert CodeMetric is not None - assert SplunkAOMetric is not None - assert LlmMetric is not None - assert LocalMetric is not None + assert AgentStream is not None + assert Evaluator is not None + assert CodeEvaluator is not None + assert SplunkAOEvaluator is not None + assert LlmEvaluator is not None + assert LocalEvaluator is not None assert MetricSpec is not None assert Model is not None assert Prompt is not None diff --git a/tests/test_base_handler.py b/tests/test_base_handler.py index 79eaadcf..58306936 100644 --- a/tests/test_base_handler.py +++ b/tests/test_base_handler.py @@ -11,7 +11,7 @@ class TestSplunkAOBaseHandler: @pytest.fixture - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def splunk_ao_logger(self, mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock): diff --git a/tests/test_crewai_handler.py b/tests/test_crewai_handler.py index c525c767..b42854c6 100644 --- a/tests/test_crewai_handler.py +++ b/tests/test_crewai_handler.py @@ -80,7 +80,7 @@ def __init__(self, raw="Test output"): def mock_splunk_ao_logger(): """Creates a mock Galileo logger for testing.""" with ( - patch("splunk_ao.logger.logger.LogStreams") as mock_logstreams, + patch("splunk_ao.logger.logger.AgentStreams") as mock_logstreams, patch("splunk_ao.logger.logger.Projects") as mock_projects, patch("splunk_ao.logger.logger.Traces") as mock_traces_client, ): diff --git a/tests/test_decorator.py b/tests/test_decorator.py index d9c9e148..b07e73e8 100644 --- a/tests/test_decorator.py +++ b/tests/test_decorator.py @@ -21,7 +21,7 @@ def reset_context(): splunk_ao_context.reset() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_context_reset( @@ -54,7 +54,7 @@ def llm_call(query: str) -> str: assert splunk_ao_context.get_current_log_stream() is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_context_init( @@ -75,7 +75,7 @@ def test_decorator_context_init( assert splunk_ao_context.get_current_log_stream() is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_context_flush( @@ -109,7 +109,7 @@ def llm_call(query: str) -> str: assert splunk_ao_context.get_current_span_stack() == [] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_context_flush_specific_project_and_log_stream( @@ -158,7 +158,7 @@ def llm_call(query: str) -> str: assert splunk_ao_context.get_current_trace() is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_context_flush_all( @@ -205,7 +205,7 @@ def llm_call(query: str) -> str: assert splunk_ao_context.get_current_span_stack() == [] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_llm_span( @@ -232,7 +232,7 @@ def llm_call(query: str) -> str: assert payload.traces[0].spans[0].output == Message(content="response", role=MessageRole.assistant) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_workflow_span_output_int( @@ -259,7 +259,7 @@ def my_function(arg1, arg2): assert payload.traces[0].spans[0].output == "3" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_workflow_span_io_object( @@ -294,7 +294,7 @@ def my_function(system: Message, user: Message): assert payload.traces[0].spans[0].output == '{"content": "response", "metadata": {"arg1": "val1", "arg2": "val2"}}' -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_tool_span_io_object( @@ -329,7 +329,7 @@ def my_function(system: Message, user: Message): assert payload.traces[0].spans[0].output == '{"content": "response", "metadata": {"arg1": "val1", "arg2": "val2"}}' -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_agent_span( @@ -356,7 +356,7 @@ def my_function(arg1: str, arg2: str) -> str: assert payload.traces[0].spans[0].output == "arg1 arg2" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_agent_span_with_agent_type( @@ -384,7 +384,7 @@ def my_function(arg1: str, arg2: str) -> str: assert payload.traces[0].spans[0].agent_type == "planner" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_agent_span_with_nested_span( @@ -421,7 +421,7 @@ def my_function(arg1: str, arg2: str): assert payload.traces[0].spans[0].spans[0].output == "arg1" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_nested_span( @@ -456,7 +456,7 @@ def nested_call(nested_query: str): assert payload.traces[0].spans[0].spans[0].output == Message(content="response", role=MessageRole.assistant) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_multiple_nested_spans( @@ -494,7 +494,7 @@ def nested_call(nested_query: str) -> str: assert payload.traces[0].spans[0].spans[0].output == Message(content="response", role=MessageRole.assistant) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_retriever_span_str( @@ -518,7 +518,7 @@ def retriever_call(query: str) -> str: assert payload.traces[0].spans[0].output == [Document(content="response1", metadata=None)] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_retriever_span_list_str( @@ -545,7 +545,7 @@ def retriever_call(query: str): ] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_retriever_span_list_dict( @@ -572,7 +572,7 @@ def retriever_call(query: str): ] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_retriever_span_list_document( @@ -599,7 +599,7 @@ def retriever_call(query: str): ] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_we_should_create_trace_but_reraise_exception( @@ -624,7 +624,7 @@ def foo() -> NoReturn: assert len(payload.traces[0].spans) == 1 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_start_session( @@ -653,7 +653,7 @@ def foo() -> str: assert payload.session_id == UUID("6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9c") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_standalone_start_session( @@ -684,7 +684,7 @@ def foo() -> str: assert payload.session_id == UUID("6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9c") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_start_session_empty_values( @@ -711,7 +711,7 @@ def foo() -> str: assert payload.session_id == UUID("6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9c") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_clear_session( @@ -742,7 +742,7 @@ def foo() -> str: assert payload.session_id is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_set_session( @@ -785,7 +785,7 @@ class ComplexPydanticModel(BaseModel): items: list = [] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_input_serialization_deserialization( @@ -816,7 +816,7 @@ def my_function(complex_input: dict) -> str: ) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_llm_span_list_output_serialization( @@ -842,7 +842,7 @@ def llm_call_returning_list(query: str): assert span.output.content == '["response1", "response2", "response3"]' -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_llm_span_tuple_output_serialization( @@ -869,7 +869,7 @@ def llm_call_returning_tuple(query: str): assert span.output.content == '["response1", "response2"]' -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_llm_span_dict_output_preserved( @@ -897,7 +897,7 @@ def llm_call_returning_dict(query: str): assert '"number": 42' in span.output.content -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_workflow_span_complex_output_serialization( @@ -928,7 +928,7 @@ def workflow_with_complex_output(query: str): assert span.output == expected_content -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_pydantic_model_input_serialization( @@ -957,7 +957,7 @@ def process_model(model: TestPydanticModel) -> str: assert '"optional_field"' not in span.input -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_pydantic_model_output_serialization( @@ -984,7 +984,7 @@ def create_model(name: str, value: int): assert '"value": 123' in span.output -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_null_output_handling( @@ -1010,7 +1010,7 @@ def function_returning_none(query: str) -> None: assert span.output is None or span.output == "" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_tool_span_output_serialization( @@ -1038,7 +1038,7 @@ def tool_with_complex_output(input_data: str): assert '"items": [1, 2, 3]' in span.output -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_agent_span_output_serialization( @@ -1066,7 +1066,7 @@ def agent_with_complex_output(query: str): assert '"actions": ["analyze", "respond"]' in span.output -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_workflow_content_blocks_output_preserved( @@ -1103,7 +1103,7 @@ def workflow_returning_content_blocks(query: str): assert isinstance(span_output, list) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_workflow_message_list_output_serialized( @@ -1138,7 +1138,7 @@ def workflow_returning_messages(query: str): # ============================================================================ -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_context_init_default( @@ -1154,7 +1154,7 @@ def test_mode_context_init_default( assert splunk_ao_context.get_current_mode() == "batch" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_context_init_explicit( @@ -1170,7 +1170,7 @@ def test_mode_context_init_explicit( assert splunk_ao_context.get_current_mode() == "distributed" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_context_call_default( @@ -1185,7 +1185,7 @@ def test_mode_context_call_default( assert splunk_ao_context.get_current_mode() == "batch" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_context_call_explicit( @@ -1200,7 +1200,7 @@ def test_mode_context_call_explicit( assert splunk_ao_context.get_current_mode() == "distributed" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_context_nested_push_pop( @@ -1223,7 +1223,7 @@ def test_mode_context_nested_push_pop( assert splunk_ao_context.get_current_mode() == "batch" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_context_multiple_nested_levels( @@ -1253,7 +1253,7 @@ def test_mode_context_multiple_nested_levels( assert splunk_ao_context.get_current_mode() == "batch" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_context_reset( @@ -1271,7 +1271,7 @@ def test_mode_context_reset( assert splunk_ao_context.get_current_mode() == "batch" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_flush_with_explicit_mode( @@ -1302,7 +1302,7 @@ def llm_call(query: str) -> str: assert splunk_ao_context.get_current_trace() is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mode_flush_different_mode_no_reset( @@ -1333,7 +1333,7 @@ def llm_call(query: str) -> str: assert splunk_ao_context.get_current_trace() == current_trace -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.dict("os.environ", {"SPLUNK_AO_MODE": "distributed"}) @@ -1351,7 +1351,7 @@ def test_mode_from_environment_variable( assert splunk_ao_context.get_current_mode() == "distributed" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.dict("os.environ", {"SPLUNK_AO_MODE": "distributed"}) @@ -1369,7 +1369,7 @@ def test_mode_explicit_overrides_environment( assert splunk_ao_context.get_current_mode() == "batch" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_get_logger_instance_with_explicit_mode( @@ -1395,7 +1395,7 @@ def test_get_logger_instance_with_explicit_mode( assert logger_distributed.mode == "distributed" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_multiple_workflow_calls_create_one_trace_with_multiple_spans( @@ -1450,7 +1450,7 @@ def process_query(query: str) -> str: assert len(logger.traces) == 0 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_session_id_context_manager( @@ -1476,7 +1476,7 @@ def foo() -> str: assert payload.session_id == UUID(test_session_id) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_session_id_nested_context_stacking( @@ -1506,7 +1506,7 @@ def test_session_id_nested_context_stacking( assert splunk_ao_context.get_logger_instance().session_id is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_session_id_cleared_on_reset_and_init( @@ -1532,7 +1532,7 @@ def test_session_id_cleared_on_reset_and_init( assert splunk_ao_context.get_logger_instance().session_id is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_start_session_overrides_context_session( @@ -1554,7 +1554,7 @@ def test_start_session_overrides_context_session( assert _session_id_context.get() == new_session_id -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_flush_on_error_called_when_flush_raises( @@ -1585,7 +1585,7 @@ def llm_call(query: str) -> str: assert isinstance(on_error.call_args[0][0], Exception) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_flush_warns_when_flush_raises_without_on_error( @@ -1613,7 +1613,7 @@ def llm_call(query: str) -> str: assert "flush failed" in mock_logger.warning.call_args[0][0] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_flush_on_error_callback_raises_is_swallowed( @@ -1643,7 +1643,7 @@ def llm_call(query: str) -> str: assert "on_error callback raised" in mock_logger.warning.call_args[0][0] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_flush_on_error_logs_at_debug_not_warning( diff --git a/tests/test_decorator_distributed.py b/tests/test_decorator_distributed.py index a14baaf4..dd408bd3 100644 --- a/tests/test_decorator_distributed.py +++ b/tests/test_decorator_distributed.py @@ -42,7 +42,7 @@ def set_distributed_mode(): os.environ["SPLUNK_AO_MODE"] = original -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_get_tracing_headers( @@ -76,7 +76,7 @@ def orchestrator(query: str) -> dict: assert headers[TRACE_ID_HEADER] == str(logger.traces[0].id) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_with_middleware_context( @@ -118,7 +118,7 @@ def downstream_service(query: str) -> str: assert logger.traces[0].name == "stub_trace" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_updates_trace_with_output_and_duration( @@ -181,7 +181,7 @@ def my_function(input_value: str) -> str: assert trace_request.is_complete, "Trace should be marked complete after flush" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_server_side_does_not_conclude_trace( @@ -229,7 +229,7 @@ def downstream_service(query: str) -> str: assert logger.traces[0].name == "stub_trace" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_client_and_server_side_behavior( @@ -320,7 +320,7 @@ def server_function(query: str) -> str: assert logger_server.traces[0].name == "stub_trace" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_workflow_span_output_is_set( @@ -370,7 +370,7 @@ def my_workflow(input_value: str) -> str: mock_traces_client_instance.update_span.assert_called() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_both_trace_and_workflow_span_have_output( @@ -436,7 +436,7 @@ def my_function(input_value: str) -> str: assert trace_request.is_complete, "Trace should be marked complete after flush" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_workflow_span_empty_string_output_is_set( @@ -480,7 +480,7 @@ def my_workflow(input_value: str) -> str: assert request.output == "", "Workflow span output should be set to empty string, not None" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_trace_duration_is_set_and_accumulates( @@ -543,7 +543,7 @@ def workflow_step_2() -> str: ) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_distributed_content_blocks_preserved_on_trace( @@ -589,7 +589,7 @@ def workflow_with_content_blocks(query: str): assert "image" in trace_request.output -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_distributed_messages_serialized_on_trace( @@ -630,7 +630,7 @@ def workflow_with_messages(query: str): assert "Hi!" in trace_request.output -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_decorator_distributed_documents_serialized_on_trace( diff --git a/tests/test_metric.py b/tests/test_evaluator.py similarity index 77% rename from tests/test_metric.py rename to tests/test_evaluator.py index 4e1b5661..82eb42f8 100644 --- a/tests/test_metric.py +++ b/tests/test_evaluator.py @@ -6,7 +6,7 @@ import pytest from galileo_core.schemas.logging.step import StepType -from splunk_ao.metric import CodeMetric, LlmMetric, LocalMetric, Metric +from splunk_ao.evaluator import CodeEvaluator, LlmEvaluator, LocalEvaluator, Evaluator from splunk_ao.resources.models import HTTPValidationError, OutputTypeEnum, ScorerTypes from splunk_ao.resources.models.invalid_result import InvalidResult from splunk_ao.resources.models.task_result_status import TaskResultStatus @@ -137,15 +137,15 @@ def _create_file(filename: str = "test_scorer.py", content: str = "def score(tra class TestMetricInitialization: - """Test suite for Metric initialization.""" + """Test suite for Evaluator initialization.""" def test_init_with_required_fields(self, reset_configuration: None) -> None: """Test initializing a metric with required fields creates a local-only instance.""" - metric = LlmMetric( - name="Test Metric", prompt="Is the response factually accurate?", model="gpt-4.1-mini", judges=3 + metric = LlmEvaluator( + name="Test Evaluator", prompt="Is the response factually accurate?", model="gpt-4.1-mini", judges=3 ) - assert metric.name == "Test Metric" + assert metric.name == "Test Evaluator" assert metric.prompt == "Is the response factually accurate?" assert metric.model == "gpt-4.1-mini" assert metric.judges == 3 @@ -155,8 +155,8 @@ def test_init_with_required_fields(self, reset_configuration: None) -> None: def test_init_with_all_fields(self, reset_configuration: None) -> None: """Test initializing a metric with all fields.""" - metric = LlmMetric( - name="Test Metric", + metric = LlmEvaluator( + name="Test Evaluator", prompt="Is the response factually accurate?", node_level=StepType.llm, cot_enabled=True, @@ -167,7 +167,7 @@ def test_init_with_all_fields(self, reset_configuration: None) -> None: output_type="percentage", ) - assert metric.name == "Test Metric" + assert metric.name == "Test Evaluator" assert metric.prompt == "Is the response factually accurate?" assert metric.scorer_type == ScorerTypes.LLM assert metric.node_level == StepType.llm @@ -180,15 +180,15 @@ def test_init_with_all_fields(self, reset_configuration: None) -> None: assert metric.ground_truth is False def test_init_ground_truth_defaults_false(self, reset_configuration: None) -> None: - # Given: an LlmMetric created without specifying ground_truth - metric = LlmMetric(name="Test Metric", prompt="Is it accurate?") + # Given: an LlmEvaluator created without specifying ground_truth + metric = LlmEvaluator(name="Test Evaluator", prompt="Is it accurate?") # Then: ground_truth defaults to False assert metric.ground_truth is False def test_init_with_ground_truth_true(self, reset_configuration: None) -> None: - # Given: an LlmMetric created with ground_truth=True - metric = LlmMetric(name="Test Metric", prompt="Compare against reference_output.", ground_truth=True) + # Given: an LlmEvaluator created with ground_truth=True + metric = LlmEvaluator(name="Test Evaluator", prompt="Compare against reference_output.", ground_truth=True) # Then: ground_truth is exposed on the instance assert metric.ground_truth is True @@ -197,19 +197,19 @@ def test_init_without_name_raises_error(self, reset_configuration: None) -> None """Test initializing a metric without name raises TypeError.""" with pytest.raises(TypeError, match="missing 1 required positional argument: 'name'"): # name is a required positional argument - omit it entirely - LlmMetric(prompt="Test prompt") # type: ignore[call-arg] + LlmEvaluator(prompt="Test prompt") # type: ignore[call-arg] def test_init_llm_scorer_without_user_prompt_raises_error(self, reset_configuration: None) -> None: """Test initializing an LLM metric without prompt raises ValidationError.""" with pytest.raises(ValidationError, match="'prompt' .* must be provided for LLM-based metrics"): - LlmMetric(name="Test Metric", prompt=None) + LlmEvaluator(name="Test Evaluator", prompt=None) class TestMetricCreate: - """Test suite for Metric.create() method.""" + """Test suite for Evaluator.create() method.""" - @patch("splunk_ao.metric.Metrics") - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Evaluators") + @patch("splunk_ao.evaluator.Scorers") def test_create_persists_metric_to_api( self, mock_scorers_class: MagicMock, mock_metrics_class: MagicMock, reset_configuration: None ) -> None: @@ -228,7 +228,7 @@ def test_create_persists_metric_to_api( mock_scorer = MagicMock() mock_scorer.id = scorer_id - mock_scorer.name = "Test Metric" + mock_scorer.name = "Test Evaluator" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = ["test"] mock_scorer.description = "Test description" @@ -242,23 +242,23 @@ def test_create_persists_metric_to_api( mock_scorer.defaults.cot_enabled = True mock_scorer.scoreable_node_types = ["llm"] - mock_metrics_service.create_custom_llm_metric.return_value = mock_version + mock_metrics_service.create_custom_llm_evaluator.return_value = mock_version mock_scorers_service.list.return_value = [mock_scorer] - metric = LlmMetric( - name="Test Metric", user_prompt="Is it accurate?", description="Test description", tags=["test"] + metric = LlmEvaluator( + name="Test Evaluator", user_prompt="Is it accurate?", description="Test description", tags=["test"] ).create() - mock_metrics_service.create_custom_llm_metric.assert_called_once() + mock_metrics_service.create_custom_llm_evaluator.assert_called_once() assert metric.id == scorer_id assert metric.is_synced() - @patch("splunk_ao.metric.Metrics") - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Evaluators") + @patch("splunk_ao.evaluator.Scorers") def test_create_forwards_ground_truth( self, mock_scorers_class: MagicMock, mock_metrics_class: MagicMock, reset_configuration: None ) -> None: - # Given: a mocked metrics service that records the create_custom_llm_metric kwargs + # Given: a mocked metrics service that records the create_custom_llm_evaluator kwargs mock_metrics_service = MagicMock() mock_metrics_class.return_value = mock_metrics_service @@ -273,7 +273,7 @@ def test_create_forwards_ground_truth( mock_scorer = MagicMock() mock_scorer.id = scorer_id - mock_scorer.name = "GT Metric" + mock_scorer.name = "GT Evaluator" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = [] mock_scorer.description = "" @@ -288,25 +288,25 @@ def test_create_forwards_ground_truth( mock_scorer.scoreable_node_types = ["llm"] mock_scorer.ground_truth = True - mock_metrics_service.create_custom_llm_metric.return_value = mock_version + mock_metrics_service.create_custom_llm_evaluator.return_value = mock_version mock_scorers_service.list.return_value = [mock_scorer] - # When: an LlmMetric with ground_truth=True is created - metric = LlmMetric(name="GT Metric", prompt="Compare against reference_output.", ground_truth=True).create() + # When: an LlmEvaluator with ground_truth=True is created + metric = LlmEvaluator(name="GT Evaluator", prompt="Compare against reference_output.", ground_truth=True).create() # Then: ground_truth=True is forwarded to the metrics service and reflected on the synced instance - _, kwargs = mock_metrics_service.create_custom_llm_metric.call_args + _, kwargs = mock_metrics_service.create_custom_llm_evaluator.call_args assert kwargs["ground_truth"] is True assert metric.ground_truth is True - @patch("splunk_ao.metric.Metrics") + @patch("splunk_ao.evaluator.Evaluators") def test_create_handles_api_failure(self, mock_metrics_class: MagicMock, reset_configuration: None) -> None: """Test create() handles API failures and sets state correctly.""" mock_service = MagicMock() mock_metrics_class.return_value = mock_service - mock_service.create_custom_llm_metric.side_effect = Exception("API Error") + mock_service.create_custom_llm_evaluator.side_effect = Exception("API Error") - metric = LlmMetric(name="Test Metric", prompt="Is it accurate?") + metric = LlmEvaluator(name="Test Evaluator", prompt="Is it accurate?") with pytest.raises(Exception, match="API Error"): metric.create() @@ -315,9 +315,9 @@ def test_create_handles_api_failure(self, mock_metrics_class: MagicMock, reset_c class TestMetricGet: - """Test suite for Metric.get() class method.""" + """Test suite for Evaluator.get() class method.""" - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_get_by_name_returns_metric(self, mock_scorers_class: MagicMock, reset_configuration: None) -> None: """Test get() with name returns a synced metric instance.""" mock_service = MagicMock() @@ -325,7 +325,7 @@ def test_get_by_name_returns_metric(self, mock_scorers_class: MagicMock, reset_c mock_scorer = MagicMock() mock_scorer.id = str(uuid4()) - mock_scorer.name = "Test Metric" + mock_scorer.name = "Test Evaluator" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = [] mock_scorer.description = "Test" @@ -341,13 +341,13 @@ def test_get_by_name_returns_metric(self, mock_scorers_class: MagicMock, reset_c mock_service.list.return_value = [mock_scorer] - metric = Metric.get(name="Test Metric") + metric = Evaluator.get(name="Test Evaluator") assert metric is not None - assert metric.name == "Test Metric" + assert metric.name == "Test Evaluator" assert metric.is_synced() - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_get_by_id_returns_metric(self, mock_scorers_class: MagicMock, reset_configuration: None) -> None: """Test get() with id returns a synced metric instance.""" mock_service = MagicMock() @@ -356,7 +356,7 @@ def test_get_by_id_returns_metric(self, mock_scorers_class: MagicMock, reset_con metric_id = str(uuid4()) mock_scorer = MagicMock() mock_scorer.id = metric_id - mock_scorer.name = "Test Metric" + mock_scorer.name = "Test Evaluator" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = [] mock_scorer.description = "Test" @@ -369,20 +369,20 @@ def test_get_by_id_returns_metric(self, mock_scorers_class: MagicMock, reset_con mock_service.list.return_value = [mock_scorer] - metric = Metric.get(id=metric_id) + metric = Evaluator.get(id=metric_id) assert metric is not None assert metric.id == metric_id assert metric.is_synced() - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_get_returns_none_when_not_found(self, mock_scorers_class: MagicMock, reset_configuration: None) -> None: """Test get() returns None when metric is not found.""" mock_service = MagicMock() mock_scorers_class.return_value = mock_service mock_service.list.return_value = [] - metric = Metric.get(name="Nonexistent Metric") + metric = Evaluator.get(name="Nonexistent Evaluator") assert metric is None @@ -396,13 +396,13 @@ def test_get_returns_none_when_not_found(self, mock_scorers_class: MagicMock, re def test_get_validates_parameters(self, kwargs: dict, expected_error: str, reset_configuration: None) -> None: """Test get() validates parameter combinations.""" with pytest.raises(ValidationError, match=expected_error): - Metric.get(**kwargs) + Evaluator.get(**kwargs) class TestMetricList: - """Test suite for Metric.list() class method.""" + """Test suite for Evaluator.list() class method.""" - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_list_returns_all_metrics(self, mock_scorers_class: MagicMock, reset_configuration: None) -> None: """Test list() returns a list of synced metric instances.""" mock_service = MagicMock() @@ -413,7 +413,7 @@ def test_list_returns_all_metrics(self, mock_scorers_class: MagicMock, reset_con for i in range(3): mock_scorer = MagicMock() mock_scorer.id = str(uuid4()) - mock_scorer.name = f"Metric {i}" + mock_scorer.name = f"Evaluator {i}" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = [] mock_scorer.description = f"Description {i}" @@ -427,13 +427,13 @@ def test_list_returns_all_metrics(self, mock_scorers_class: MagicMock, reset_con mock_service.list.return_value = mock_scorers - metrics = Metric.list() + metrics = Evaluator.list() assert len(metrics) == 3 - assert all(isinstance(m, Metric) for m in metrics) + assert all(isinstance(m, Evaluator) for m in metrics) assert all(m.is_synced() for m in metrics) - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_list_with_name_filter(self, mock_scorers_class: MagicMock, reset_configuration: None) -> None: """Test list() with name filter.""" mock_service = MagicMock() @@ -441,7 +441,7 @@ def test_list_with_name_filter(self, mock_scorers_class: MagicMock, reset_config mock_scorer = MagicMock() mock_scorer.id = str(uuid4()) - mock_scorer.name = "Factuality Metric" + mock_scorer.name = "Factuality Evaluator" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = [] mock_scorer.description = "Test" @@ -454,29 +454,29 @@ def test_list_with_name_filter(self, mock_scorers_class: MagicMock, reset_config mock_service.list.return_value = [mock_scorer] - metrics = Metric.list(name_filter="Factuality Metric") + metrics = Evaluator.list(name_filter="Factuality Evaluator") - mock_service.list.assert_called_once_with(name="Factuality Metric", types=None) + mock_service.list.assert_called_once_with(name="Factuality Evaluator", types=None) assert len(metrics) == 1 - assert metrics[0].name == "Factuality Metric" + assert metrics[0].name == "Factuality Evaluator" - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_list_with_scorer_types_filter(self, mock_scorers_class: MagicMock, reset_configuration: None) -> None: """Test list() with scorer types filter.""" mock_service = MagicMock() mock_scorers_class.return_value = mock_service mock_service.list.return_value = [] - Metric.list(scorer_types=[ScorerTypes.LLM]) + Evaluator.list(scorer_types=[ScorerTypes.LLM]) mock_service.list.assert_called_once_with(name=None, types=[ScorerTypes.LLM]) class TestMetricDelete: - """Test suite for Metric.delete() method.""" + """Test suite for Evaluator.delete() method.""" - @patch("splunk_ao.metric.Metrics") - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Evaluators") + @patch("splunk_ao.evaluator.Scorers") def test_delete_removes_metric( self, mock_scorers_class: MagicMock, mock_metrics_class: MagicMock, reset_configuration: None ) -> None: @@ -490,7 +490,7 @@ def test_delete_removes_metric( metric_id = str(uuid4()) mock_scorer = MagicMock() mock_scorer.id = metric_id - mock_scorer.name = "Test Metric" + mock_scorer.name = "Test Evaluator" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = [] mock_scorer.description = "Test" @@ -503,24 +503,24 @@ def test_delete_removes_metric( mock_scorers_service.list.return_value = [mock_scorer] - metric = Metric.get(id=metric_id) + metric = Evaluator.get(id=metric_id) metric.delete() - mock_metrics_service.delete_metric.assert_called_once_with(name="Test Metric") + mock_metrics_service.delete_evaluator.assert_called_once_with(name="Test Evaluator") assert metric.sync_state == SyncState.DELETED def test_delete_raises_error_for_local_only(self, reset_configuration: None) -> None: """Test delete() raises ValueError for local-only metric.""" - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") - with pytest.raises(ValueError, match="Metric ID is not set"): + with pytest.raises(ValueError, match="Evaluator ID is not set"): metric.delete() class TestMetricRefresh: - """Test suite for Metric.refresh() method.""" + """Test suite for Evaluator.refresh() method.""" - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_refresh_updates_attributes(self, mock_scorers_class: MagicMock, reset_configuration: None) -> None: """Test refresh() updates all attributes from the API.""" mock_service = MagicMock() @@ -531,7 +531,7 @@ def test_refresh_updates_attributes(self, mock_scorers_class: MagicMock, reset_c # Initial state initial_scorer = MagicMock() initial_scorer.id = metric_id - initial_scorer.name = "Test Metric" + initial_scorer.name = "Test Evaluator" initial_scorer.scorer_type = ScorerTypes.LLM initial_scorer.tags = ["test"] initial_scorer.description = "Initial description" @@ -548,7 +548,7 @@ def test_refresh_updates_attributes(self, mock_scorers_class: MagicMock, reset_c # Updated state updated_scorer = MagicMock() updated_scorer.id = metric_id - updated_scorer.name = "Test Metric" + updated_scorer.name = "Test Evaluator" updated_scorer.scorer_type = ScorerTypes.LLM updated_scorer.tags = ["test", "updated"] updated_scorer.description = "Updated description" @@ -564,7 +564,7 @@ def test_refresh_updates_attributes(self, mock_scorers_class: MagicMock, reset_c mock_service.list.side_effect = [[initial_scorer], [updated_scorer]] - metric = Metric.get(id=metric_id) + metric = Evaluator.get(id=metric_id) assert metric.description == "Initial description" assert metric.judges == 3 @@ -576,12 +576,12 @@ def test_refresh_updates_attributes(self, mock_scorers_class: MagicMock, reset_c def test_refresh_raises_error_for_local_only(self, reset_configuration: None) -> None: """Test refresh() raises ValueError for local-only metric.""" - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") - with pytest.raises(ValueError, match="Metric ID is not set"): + with pytest.raises(ValueError, match="Evaluator ID is not set"): metric.refresh() - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_refresh_raises_error_when_metric_no_longer_exists( self, mock_scorers_class: MagicMock, reset_configuration: None ) -> None: @@ -592,7 +592,7 @@ def test_refresh_raises_error_when_metric_no_longer_exists( metric_id = str(uuid4()) mock_scorer = MagicMock() mock_scorer.id = metric_id - mock_scorer.name = "Test Metric" + mock_scorer.name = "Test Evaluator" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = [] mock_scorer.description = "Test" @@ -606,21 +606,21 @@ def test_refresh_raises_error_when_metric_no_longer_exists( # First call returns the metric, second returns empty list (deleted) mock_service.list.side_effect = [[mock_scorer], []] - metric = Metric.get(id=metric_id) + metric = Evaluator.get(id=metric_id) with pytest.raises(ValueError, match="no longer exists"): metric.refresh() class TestMetricUpdate: - """Test suite for Metric.update() method.""" + """Test suite for Evaluator.update() method.""" def test_update_local_metric_raises_validation_error(self, reset_configuration: None) -> None: # Given: a local metric instance def dummy_fn(trace): return 0.5 - metric = LocalMetric(name="my-local", scorer_fn=dummy_fn) + metric = LocalEvaluator(name="my-local", scorer_fn=dummy_fn) # When/Then: calling update raises ValidationError with pytest.raises(ValidationError, match="Local metrics don't exist"): @@ -628,15 +628,15 @@ def dummy_fn(trace): def test_update_without_id_raises_value_error(self, reset_configuration: None) -> None: # Given: an LLM metric without an ID (local-only) - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") # When/Then: calling update raises ValueError about missing ID - with pytest.raises(ValueError, match="Metric ID is not set"): + with pytest.raises(ValueError, match="Evaluator ID is not set"): metric.update(name="New Name") def test_update_deleted_metric_raises_value_error(self, reset_configuration: None) -> None: # Given: a metric in DELETED state - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") metric._sync_attrs(id="some-id") metric._set_state(SyncState.DELETED) @@ -646,7 +646,7 @@ def test_update_deleted_metric_raises_value_error(self, reset_configuration: Non def test_update_failed_sync_raises_value_error(self, reset_configuration: None) -> None: # Given: a metric in FAILED_SYNC state - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") metric._sync_attrs(id="some-id") metric._set_state(SyncState.FAILED_SYNC, error=RuntimeError("prior failure")) @@ -656,7 +656,7 @@ def test_update_failed_sync_raises_value_error(self, reset_configuration: None) def test_update_with_invalid_fields_raises_value_error(self, reset_configuration: None) -> None: # Given: a synced LLM metric - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") metric._sync_attrs(id="some-id") metric._set_state(SyncState.SYNCED) @@ -664,14 +664,14 @@ def test_update_with_invalid_fields_raises_value_error(self, reset_configuration with pytest.raises(ValueError, match="Invalid update fields"): metric.update(prompt="new prompt", model="gpt-4o") - @patch("splunk_ao.metric.update_scorers_scorer_id_patch") - @patch("splunk_ao.metric.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.update_scorers_scorer_id_patch") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") def test_update_calls_api_and_syncs_attributes( self, mock_config_get: MagicMock, mock_update_patch: MagicMock, reset_configuration: None ) -> None: # Given: a synced LLM metric with an ID metric_id = str(uuid4()) - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") metric._sync_attrs(id=metric_id, tags=["old-tag"], description="old desc") metric._set_state(SyncState.SYNCED) @@ -680,7 +680,7 @@ def test_update_calls_api_and_syncs_attributes( updated_response = MagicMock() updated_response.id = metric_id - updated_response.name = "Test Metric" + updated_response.name = "Test Evaluator" updated_response.scorer_type = ScorerTypes.LLM updated_response.tags = ["new-tag"] updated_response.description = "new desc" @@ -705,13 +705,13 @@ def test_update_calls_api_and_syncs_attributes( assert metric.tags == ["new-tag"] assert metric.sync_state == SyncState.SYNCED - @patch("splunk_ao.metric.update_scorers_scorer_id_patch") - @patch("splunk_ao.metric.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.update_scorers_scorer_id_patch") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") def test_update_handles_api_failure( self, mock_config_get: MagicMock, mock_update_patch: MagicMock, reset_configuration: None ) -> None: # Given: a synced metric whose API call will fail - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") metric._sync_attrs(id="some-id") metric._set_state(SyncState.SYNCED) @@ -724,13 +724,13 @@ def test_update_handles_api_failure( assert metric.sync_state == SyncState.FAILED_SYNC - @patch("splunk_ao.metric.update_scorers_scorer_id_patch") - @patch("splunk_ao.metric.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.update_scorers_scorer_id_patch") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") def test_update_raises_api_error_for_validation_error_response( self, mock_config_get: MagicMock, mock_update_patch: MagicMock, reset_configuration: None ) -> None: # Given: a synced metric and an API that returns an HTTPValidationError - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") metric._sync_attrs(id="some-id") metric._set_state(SyncState.SYNCED) @@ -740,16 +740,16 @@ def test_update_raises_api_error_for_validation_error_response( mock_update_patch.sync.return_value = validation_error_response # When/Then: APIError is raised with the validation detail - with pytest.raises(APIError, match="Metric update validation error"): + with pytest.raises(APIError, match="Evaluator update validation error"): metric.update(name="New Name") - @patch("splunk_ao.metric.update_scorers_scorer_id_patch") - @patch("splunk_ao.metric.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.update_scorers_scorer_id_patch") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") def test_update_raises_api_error_for_none_response( self, mock_config_get: MagicMock, mock_update_patch: MagicMock, reset_configuration: None ) -> None: # Given: a synced metric and an API that returns None - metric = LlmMetric(name="Test Metric", prompt="Test prompt") + metric = LlmEvaluator(name="Test Evaluator", prompt="Test prompt") metric._sync_attrs(id="some-id") metric._set_state(SyncState.SYNCED) @@ -762,17 +762,17 @@ def test_update_raises_api_error_for_none_response( class TestMetricMethods: - """Test suite for other Metric methods.""" + """Test suite for other Evaluator methods.""" def test_str_and_repr(self, reset_configuration: None) -> None: """Test __str__ and __repr__ return expected formats.""" - metric = LlmMetric(name="Test Metric", prompt="Is it accurate?", output_type="percentage") + metric = LlmEvaluator(name="Test Evaluator", prompt="Is it accurate?", output_type="percentage") metric.id = "test-id-123" - assert str(metric) == "LlmMetric(name='Test Metric', id='test-id-123', scorer_type='llm')" + assert str(metric) == "LlmEvaluator(name='Test Evaluator', id='test-id-123', scorer_type='llm')" assert "model=" in repr(metric) and "judges=" in repr(metric) - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.Scorers") def test_populate_from_scorer_response_handles_unset_values( self, mock_scorers_class: MagicMock, reset_configuration: None ) -> None: @@ -784,7 +784,7 @@ def test_populate_from_scorer_response_handles_unset_values( mock_scorer = MagicMock() mock_scorer.id = str(uuid4()) - mock_scorer.name = "Test Metric" + mock_scorer.name = "Test Evaluator" mock_scorer.scorer_type = ScorerTypes.LLM mock_scorer.tags = [] mock_scorer.description = UnsetType() @@ -797,7 +797,7 @@ def test_populate_from_scorer_response_handles_unset_values( mock_service.list.return_value = [mock_scorer] - metric = Metric.get(name="Test Metric") + metric = Evaluator.get(name="Test Evaluator") assert metric is not None assert metric.description == "" @@ -812,43 +812,43 @@ def test_populate_from_scorer_response_handles_unset_values( class TestCodeMetricInitialization: - """Test suite for CodeMetric initialization.""" + """Test suite for CodeEvaluator initialization.""" def test_init_with_minimal_fields(self, reset_configuration: None) -> None: - """Test initializing a CodeMetric with just a name.""" - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + """Test initializing a CodeEvaluator with just a name.""" + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) - assert metric.name == "Test Code Metric" + assert metric.name == "Test Code Evaluator" assert metric.scorer_type == ScorerTypes.CODE assert metric.id is None assert metric.sync_state == SyncState.LOCAL_ONLY def test_init_with_all_fields(self, reset_configuration: None) -> None: - """Test initializing a CodeMetric with all optional fields.""" - metric = CodeMetric(name="Test Code Metric", description="Test code metric description", tags=["test", "code"]) + """Test initializing a CodeEvaluator with all optional fields.""" + metric = CodeEvaluator(name="Test Code Evaluator", description="Test code metric description", tags=["test", "code"]) - assert metric.name == "Test Code Metric" + assert metric.name == "Test Code Evaluator" assert metric.description == "Test code metric description" assert metric.tags == ["test", "code"] def test_init_with_required_metrics(self, reset_configuration: None) -> None: - """Test initializing a CodeMetric with required_metrics.""" - metric = CodeMetric( - name="Test Code Metric", node_level=StepType.llm, required_metrics=["context_adherence", "completeness"] + """Test initializing a CodeEvaluator with required_metrics.""" + metric = CodeEvaluator( + name="Test Code Evaluator", node_level=StepType.llm, required_metrics=["context_adherence", "completeness"] ) - assert metric.name == "Test Code Metric" + assert metric.name == "Test Code Evaluator" assert metric.required_metrics == ["context_adherence", "completeness"] assert metric.node_level == StepType.llm def test_init_with_output_type_enum(self, reset_configuration: None) -> None: - """Test initializing a CodeMetric with an OutputTypeEnum output_type.""" - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm, output_type=OutputTypeEnum.PERCENTAGE) + """Test initializing a CodeEvaluator with an OutputTypeEnum output_type.""" + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm, output_type=OutputTypeEnum.PERCENTAGE) assert metric.output_type == OutputTypeEnum.PERCENTAGE def test_init_with_output_type_string(self, reset_configuration: None) -> None: - """Test initializing a CodeMetric with a string output_type is mapped to the enum.""" + """Test initializing a CodeEvaluator with a string output_type is mapped to the enum.""" cases = [ ("percentage", OutputTypeEnum.PERCENTAGE), ("boolean", OutputTypeEnum.BOOLEAN), @@ -857,25 +857,25 @@ def test_init_with_output_type_string(self, reset_configuration: None) -> None: ("discrete", OutputTypeEnum.DISCRETE), ] for string_value, expected_enum in cases: - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm, output_type=string_value) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm, output_type=string_value) assert metric.output_type == expected_enum, f"Expected {expected_enum} for '{string_value}'" def test_init_without_output_type_defaults_to_none(self, reset_configuration: None) -> None: """Test that omitting output_type leaves it as None (API applies its own default).""" - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) assert metric.output_type is None class TestCodeMetricCreate: - """Test suite for CodeMetric.create() method.""" - - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post") - @patch("splunk_ao.metric.create_scorers_post") - @patch("splunk_ao.metric.Scorers") + """Test suite for CodeEvaluator.create() method.""" + + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post") + @patch("splunk_ao.evaluator.create_scorers_post") + @patch("splunk_ao.evaluator.Scorers") def test_create_persists_code_metric_to_api( self, mock_scorers_class: MagicMock, @@ -907,7 +907,7 @@ def test_create_persists_code_metric_to_api( # Mock the scorer creation response scorer_id = str(uuid4()) - mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Metric") + mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Evaluator") # Mock the version creation response mock_create_version.sync.return_value = mock_version_response(scorer_id) @@ -915,13 +915,13 @@ def test_create_persists_code_metric_to_api( # Mock the scorer list response for refresh mock_scorers_service = MagicMock() mock_scorers_service.list.return_value = [ - mock_scorer_full(scorer_id, "Test Code Metric", tags=["test"], description="Test description") + mock_scorer_full(scorer_id, "Test Code Evaluator", tags=["test"], description="Test description") ] mock_scorers_class.return_value = mock_scorers_service # Create the metric metric = ( - CodeMetric(name="Test Code Metric", description="Test description", tags=["test"], node_level=StepType.llm) + CodeEvaluator(name="Test Code Evaluator", description="Test description", tags=["test"], node_level=StepType.llm) .load_code(str(code_file)) .create() ) @@ -933,7 +933,7 @@ def test_create_persists_code_metric_to_api( # Verify scorer creation was called mock_create_scorers.sync.assert_called_once() create_scorer_call = mock_create_scorers.sync.call_args - assert create_scorer_call.kwargs["body"].name == "Test Code Metric" + assert create_scorer_call.kwargs["body"].name == "Test Code Evaluator" assert create_scorer_call.kwargs["body"].scorer_type == ScorerTypes.CODE assert create_scorer_call.kwargs["body"].description == "Test description" assert create_scorer_call.kwargs["body"].tags == ["test"] @@ -954,12 +954,12 @@ def test_create_persists_code_metric_to_api( assert metric.id == scorer_id assert metric.is_synced() - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post") - @patch("splunk_ao.metric.create_scorers_post") - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post") + @patch("splunk_ao.evaluator.create_scorers_post") + @patch("splunk_ao.evaluator.Scorers") def test_create_forwards_output_type_to_scorer_request( self, mock_scorers_class: MagicMock, @@ -985,23 +985,23 @@ def test_create_forwards_output_type_to_scorer_request( mock_validate_post.sync.return_value = mock_validation_response(task_id) mock_validate_get.sync.return_value = mock_validation_task_result(TaskResultStatus.COMPLETED) scorer_id = str(uuid4()) - mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Metric") + mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Evaluator") mock_create_version.sync.return_value = mock_version_response(scorer_id) mock_scorers_service = MagicMock() - mock_scorers_service.list.return_value = [mock_scorer_full(scorer_id, "Test Code Metric")] + mock_scorers_service.list.return_value = [mock_scorer_full(scorer_id, "Test Code Evaluator")] mock_scorers_class.return_value = mock_scorers_service - CodeMetric(name="Test Code Metric", node_level=StepType.llm, output_type=OutputTypeEnum.PERCENTAGE).load_code( + CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm, output_type=OutputTypeEnum.PERCENTAGE).load_code( str(code_file) ).create() create_scorer_call = mock_create_scorers.sync.call_args assert create_scorer_call.kwargs["body"].output_type == OutputTypeEnum.PERCENTAGE - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_scorers_post") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_scorers_post") def test_create_handles_scorer_creation_failure( self, mock_create_scorers: MagicMock, @@ -1028,18 +1028,18 @@ def test_create_handles_scorer_creation_failure( # Mock scorer creation to fail mock_create_scorers.sync.side_effect = Exception("Scorer creation failed") - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) with pytest.raises(Exception, match="Scorer creation failed"): metric.load_code(str(code_file)).create() assert metric.sync_state == SyncState.FAILED_SYNC - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post") - @patch("splunk_ao.metric.create_scorers_post") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post") + @patch("splunk_ao.evaluator.create_scorers_post") def test_create_handles_version_creation_failure( self, mock_create_scorers: MagicMock, @@ -1067,24 +1067,24 @@ def test_create_handles_version_creation_failure( # Mock scorer creation to succeed scorer_id = str(uuid4()) - mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Metric") + mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Evaluator") # Mock version creation to fail mock_create_version.sync.side_effect = Exception("Version creation failed") - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) with pytest.raises(Exception, match="Version creation failed"): metric.load_code(str(code_file)).create() assert metric.sync_state == SyncState.FAILED_SYNC - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post") - @patch("splunk_ao.metric.create_scorers_post") - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post") + @patch("splunk_ao.evaluator.create_scorers_post") + @patch("splunk_ao.evaluator.Scorers") def test_create_with_different_node_levels( self, mock_scorers_class: MagicMock, @@ -1116,7 +1116,7 @@ def test_create_with_different_node_levels( # Mock the scorer creation response scorer_id = str(uuid4()) - mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, f"Test Code Metric {node_level}") + mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, f"Test Code Evaluator {node_level}") # Mock the version creation response mock_create_version.sync.return_value = mock_version_response(scorer_id) @@ -1124,13 +1124,13 @@ def test_create_with_different_node_levels( # Mock the scorer list response for refresh mock_scorers_service = MagicMock() mock_scorers_service.list.return_value = [ - mock_scorer_full(scorer_id, f"Test Code Metric {node_level}", node_types=[node_level]) + mock_scorer_full(scorer_id, f"Test Code Evaluator {node_level}", node_types=[node_level]) ] mock_scorers_class.return_value = mock_scorers_service # Create the metric metric = ( - CodeMetric(name=f"Test Code Metric {node_level}", node_level=node_level) + CodeEvaluator(name=f"Test Code Evaluator {node_level}", node_level=node_level) .load_code(str(code_file)) .create() ) @@ -1140,12 +1140,12 @@ def test_create_with_different_node_levels( # Verify the node_level is set on the metric itself assert metric.node_level == node_level - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post") - @patch("splunk_ao.metric.create_scorers_post") - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post") + @patch("splunk_ao.evaluator.create_scorers_post") + @patch("splunk_ao.evaluator.Scorers") def test_create_with_required_metrics( self, mock_scorers_class: MagicMock, @@ -1177,21 +1177,21 @@ def test_create_with_required_metrics( # Mock the scorer creation response scorer_id = str(uuid4()) - mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Metric") + mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Evaluator") # Mock the version creation response mock_create_version.sync.return_value = mock_version_response(scorer_id) # Mock the scorer list response for refresh mock_scorers_service = MagicMock() - mock_scorers_service.list.return_value = [mock_scorer_full(scorer_id, "Test Code Metric")] + mock_scorers_service.list.return_value = [mock_scorer_full(scorer_id, "Test Code Evaluator")] mock_scorers_class.return_value = mock_scorers_service required_metrics = ["context_adherence", "completeness"] # Create the metric with required_metrics metric = ( - CodeMetric(name="Test Code Metric", node_level=StepType.llm, required_metrics=required_metrics) + CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm, required_metrics=required_metrics) .load_code(str(code_file)) .create() ) @@ -1207,7 +1207,7 @@ def test_create_with_required_metrics( create_scorer_call = mock_create_scorers.sync.call_args assert create_scorer_call.kwargs["body"].required_scorers == required_metrics - @patch("splunk_ao.metric.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") def test_create_reads_code_file_correctly( self, mock_config: MagicMock, @@ -1234,15 +1234,15 @@ def score(trace): code_file = create_temp_code_file(filename="complex_scorer.py", content=expected_content) with ( - patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") as mock_validate_post, + patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") as mock_validate_post, patch( - "splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get" + "splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get" ) as mock_validate_get, - patch("splunk_ao.metric.create_scorers_post") as mock_create_scorers, + patch("splunk_ao.evaluator.create_scorers_post") as mock_create_scorers, patch( - "splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post" + "splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post" ) as mock_create_version, - patch("splunk_ao.metric.Scorers") as mock_scorers_class, + patch("splunk_ao.evaluator.Scorers") as mock_scorers_class, ): # Mock validation flow mock_validate_post.sync.return_value = mock_validation_response() @@ -1258,7 +1258,7 @@ def score(trace): mock_scorers_class.return_value = mock_scorers_service # Create the metric - CodeMetric(name="Complex Scorer", node_level=StepType.llm).load_code(str(code_file)).create() + CodeEvaluator(name="Complex Scorer", node_level=StepType.llm).load_code(str(code_file)).create() # Verify the code file was read as bytes version_call = mock_create_version.sync.call_args @@ -1267,8 +1267,8 @@ def score(trace): assert hasattr(body, "file") assert hasattr(body.file, "payload") - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.create_scorers_post") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.create_scorers_post") def test_load_code_with_nonexistent_file_raises_validation_error( self, mock_create_scorers: MagicMock, @@ -1282,15 +1282,15 @@ def test_load_code_with_nonexistent_file_raises_validation_error( mock_config.return_value.api_client = mock_api_client # Create a metric and pass a non-existent file path - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) with pytest.raises(ValidationError, match="Code file not found"): metric.load_code("/nonexistent/file.py").create() - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_scorers_post") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_scorers_post") def test_create_handles_none_scorer_response( self, mock_create_scorers: MagicMock, @@ -1317,18 +1317,18 @@ def test_create_handles_none_scorer_response( # Mock scorer creation to return None mock_create_scorers.sync.return_value = None - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) with pytest.raises(ValueError, match="Failed to create code-based metric: No response from API"): metric.load_code(str(code_file)).create() assert metric.sync_state == SyncState.FAILED_SYNC - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post") - @patch("splunk_ao.metric.create_scorers_post") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post") + @patch("splunk_ao.evaluator.create_scorers_post") def test_create_handles_none_version_response( self, mock_create_scorers: MagicMock, @@ -1356,22 +1356,22 @@ def test_create_handles_none_version_response( # Mock scorer creation to succeed scorer_id = str(uuid4()) - mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Metric") + mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Evaluator") # Mock version creation to return None mock_create_version.sync.return_value = None - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) with pytest.raises(ValueError, match="Failed to create code-based metric: No response from API"): metric.load_code(str(code_file)).create() assert metric.sync_state == SyncState.FAILED_SYNC - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_scorers_post") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_scorers_post") def test_create_propagates_validation_error( self, mock_create_scorers: MagicMock, @@ -1399,16 +1399,16 @@ def test_create_propagates_validation_error( validation_error = ValidationError("Invalid configuration") mock_create_scorers.sync.side_effect = validation_error - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) # ValidationError should be propagated as-is, not wrapped with pytest.raises(ValidationError, match="Invalid configuration"): metric.load_code(str(code_file)).create() - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_scorers_post") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_scorers_post") def test_create_sets_failed_sync_state_on_general_exception( self, mock_create_scorers: MagicMock, @@ -1436,7 +1436,7 @@ def test_create_sets_failed_sync_state_on_general_exception( runtime_error = RuntimeError("Unexpected error") mock_create_scorers.sync.side_effect = runtime_error - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) with pytest.raises(RuntimeError, match="Unexpected error"): metric.load_code(str(code_file)).create() @@ -1445,9 +1445,9 @@ def test_create_sets_failed_sync_state_on_general_exception( assert metric.sync_state == SyncState.FAILED_SYNC assert metric._last_error == runtime_error - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") def test_create_handles_validation_failure( self, mock_validate_post: MagicMock, @@ -1470,18 +1470,18 @@ def test_create_handles_validation_failure( mock_validate_post.sync.return_value = mock_validation_response() mock_validate_get.sync.return_value = mock_validation_task_result(TaskResultStatus.FAILED) - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm) + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm) with pytest.raises(ValidationError, match="Code validation failed"): metric.load_code(str(code_file)).create() - @patch("splunk_ao.metric.time.sleep") - @patch("splunk_ao.metric.SplunkAOConfig.get") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post") - @patch("splunk_ao.metric.create_scorers_post") - @patch("splunk_ao.metric.Scorers") + @patch("splunk_ao.evaluator.time.sleep") + @patch("splunk_ao.evaluator.SplunkAOConfig.get") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post") + @patch("splunk_ao.evaluator.create_scorers_post") + @patch("splunk_ao.evaluator.Scorers") def test_create_polls_until_validation_complete( self, mock_scorers_class: MagicMock, @@ -1517,16 +1517,16 @@ def test_create_polls_until_validation_complete( # Mock the scorer creation response scorer_id = str(uuid4()) - mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Metric") + mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test Code Evaluator") mock_create_version.sync.return_value = mock_version_response(scorer_id) # Mock the scorer list response for refresh mock_scorers_service = MagicMock() - mock_scorers_service.list.return_value = [mock_scorer_full(scorer_id, "Test Code Metric")] + mock_scorers_service.list.return_value = [mock_scorer_full(scorer_id, "Test Code Evaluator")] mock_scorers_class.return_value = mock_scorers_service # Create the metric - metric = CodeMetric(name="Test Code Metric", node_level=StepType.llm).load_code(str(code_file)).create() + metric = CodeEvaluator(name="Test Code Evaluator", node_level=StepType.llm).load_code(str(code_file)).create() # Verify polling was called 3 times (2 pending + 1 completed) assert mock_validate_get.sync.call_count == 3 @@ -1534,11 +1534,11 @@ def test_create_polls_until_validation_complete( assert mock_sleep.call_count == 2 assert metric.is_synced() - @patch("splunk_ao.metric.time.time") - @patch("splunk_ao.metric.time.sleep") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.time.time") + @patch("splunk_ao.evaluator.time.sleep") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_create_validation_timeout( self, mock_config, @@ -1563,14 +1563,14 @@ def test_create_validation_timeout( # Simulate time passing: first call is start_time=0, second call is elapsed=61 (past 60s timeout) mock_time.side_effect = [0.0, 61.0] - metric = CodeMetric(name="Test Timeout Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Timeout Evaluator", node_level=StepType.llm).load_code(str(code_file)) with pytest.raises(ValidationError, match="Code validation timed out"): metric.create() - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_create_validation_post_returns_none( self, mock_config, mock_validate_post, mock_validate_get, create_temp_code_file, mock_api_client ) -> None: @@ -1581,14 +1581,14 @@ def test_create_validation_post_returns_none( mock_validate_post.sync.return_value = None - metric = CodeMetric(name="Test None Response Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test None Response Evaluator", node_level=StepType.llm).load_code(str(code_file)) with pytest.raises(ValueError, match="Failed to validate code: No response from API"): metric.create() - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_create_validation_get_returns_none( self, mock_config, @@ -1606,14 +1606,14 @@ def test_create_validation_get_returns_none( mock_validate_post.sync.return_value = mock_validation_response() mock_validate_get.sync.return_value = None - metric = CodeMetric(name="Test None Get Response Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test None Get Response Evaluator", node_level=StepType.llm).load_code(str(code_file)) with pytest.raises(ValueError, match="Failed to get validation result: No response from API"): metric.create() - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_create_validation_unknown_status( self, mock_config, @@ -1634,14 +1634,14 @@ def test_create_validation_unknown_status( mock_result.status = "unknown_status" mock_validate_get.sync.return_value = mock_result - metric = CodeMetric(name="Test Unknown Status Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Unknown Status Evaluator", node_level=StepType.llm).load_code(str(code_file)) with pytest.raises(ValueError, match="Unknown task status"): metric.create() - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_create_validation_failed_status( self, mock_config, @@ -1662,14 +1662,14 @@ def test_create_validation_failed_status( mock_result.result = "Syntax error in code" mock_validate_get.sync.return_value = mock_result - metric = CodeMetric(name="Test Failed Status Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Failed Status Evaluator", node_level=StepType.llm).load_code(str(code_file)) with pytest.raises(ValidationError, match="Code validation failed: Syntax error in code"): metric.create() - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_create_validation_failed_status_no_message( self, mock_config, @@ -1690,14 +1690,14 @@ def test_create_validation_failed_status_no_message( mock_result.result = None # No error message mock_validate_get.sync.return_value = mock_result - metric = CodeMetric(name="Test Failed No Message Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Failed No Message Evaluator", node_level=StepType.llm).load_code(str(code_file)) with pytest.raises(ValidationError, match="Code validation failed"): metric.create() - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_create_validation_invalid_result( self, mock_config, @@ -1734,17 +1734,17 @@ def test_create_validation_invalid_result( mock_task_result.result = mock_validate_result mock_validate_get.sync.return_value = mock_task_result - metric = CodeMetric(name="Test Invalid Result Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Invalid Result Evaluator", node_level=StepType.llm).load_code(str(code_file)) with pytest.raises(ValidationError, match="Code validation failed: Missing required function: evaluate"): metric.create() - @patch("splunk_ao.metric.Scorers") - @patch("splunk_ao.metric.create_code_scorer_version_scorers_scorer_id_version_code_post") - @patch("splunk_ao.metric.create_scorers_post") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.Scorers") + @patch("splunk_ao.evaluator.create_code_scorer_version_scorers_scorer_id_version_code_post") + @patch("splunk_ao.evaluator.create_scorers_post") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_create_validation_result_as_string( self, mock_config, @@ -1774,15 +1774,15 @@ def test_create_validation_result_as_string( mock_validate_get.sync.return_value = mock_task_result scorer_id = str(uuid4()) - mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test String Result Metric") + mock_create_scorers.sync.return_value = mock_scorer_response(scorer_id, "Test String Result Evaluator") mock_create_version.sync.return_value = mock_version_response(scorer_id) mock_scorers_service = MagicMock() - mock_scorers_service.list.return_value = [mock_scorer_full(scorer_id, "Test String Result Metric")] + mock_scorers_service.list.return_value = [mock_scorer_full(scorer_id, "Test String Result Evaluator")] mock_scorers_class.return_value = mock_scorers_service metric = ( - CodeMetric(name="Test String Result Metric", node_level=StepType.llm).load_code(str(code_file)).create() + CodeEvaluator(name="Test String Result Evaluator", node_level=StepType.llm).load_code(str(code_file)).create() ) assert metric.is_synced() @@ -1790,13 +1790,13 @@ def test_create_validation_result_as_string( class TestCodeMetricValidationConfiguration: - """Test suite for CodeMetric validation configuration parameters.""" + """Test suite for CodeEvaluator validation configuration parameters.""" - @patch("splunk_ao.metric.time.time") - @patch("splunk_ao.metric.time.sleep") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.time.time") + @patch("splunk_ao.evaluator.time.sleep") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_custom_timeout_value_is_respected( self, mock_config, @@ -1825,16 +1825,16 @@ def test_custom_timeout_value_is_respected( # Simulate time passing: first call is start_time=0, second call is elapsed=31 (past 30s custom timeout) mock_time.side_effect = [0.0, 31.0] - metric = CodeMetric(name="Test Custom Timeout Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Custom Timeout Evaluator", node_level=StepType.llm).load_code(str(code_file)) with pytest.raises(ValidationError, match="Code validation timed out after 30 seconds"): metric.create() - @patch("splunk_ao.metric.time.time") - @patch("splunk_ao.metric.time.sleep") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.time.time") + @patch("splunk_ao.evaluator.time.sleep") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_custom_initial_delay_is_used( self, mock_config, @@ -1868,7 +1868,7 @@ def test_custom_initial_delay_is_used( # Time progression mock_time.side_effect = [0.0, 1.0, 2.0] - metric = CodeMetric(name="Test Initial Delay Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Initial Delay Evaluator", node_level=StepType.llm).load_code(str(code_file)) # Should raise because no scorers mock, but we can check the sleep was called correctly try: @@ -1881,11 +1881,11 @@ def test_custom_initial_delay_is_used( first_sleep_call = mock_sleep.call_args_list[0] assert first_sleep_call[0][0] == 2.0 - @patch("splunk_ao.metric.time.time") - @patch("splunk_ao.metric.time.sleep") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.time.time") + @patch("splunk_ao.evaluator.time.sleep") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_custom_max_delay_caps_backoff( self, mock_config, @@ -1922,7 +1922,7 @@ def test_custom_max_delay_caps_backoff( # Time progression mock_time.side_effect = [0.0, 1.0, 2.0, 3.0, 4.0] - metric = CodeMetric(name="Test Max Delay Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Max Delay Evaluator", node_level=StepType.llm).load_code(str(code_file)) try: metric.create() @@ -1934,11 +1934,11 @@ def test_custom_max_delay_caps_backoff( second_sleep_call = mock_sleep.call_args_list[1] assert second_sleep_call[0][0] == 15.0 # Capped at max_delay - @patch("splunk_ao.metric.time.time") - @patch("splunk_ao.metric.time.sleep") - @patch("splunk_ao.metric.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") - @patch("splunk_ao.metric.validate_code_scorer_scorers_code_validate_post") - @patch("splunk_ao.metric.SplunkAOConfig") + @patch("splunk_ao.evaluator.time.time") + @patch("splunk_ao.evaluator.time.sleep") + @patch("splunk_ao.evaluator.get_validate_code_scorer_task_result_scorers_code_validate_task_id_get") + @patch("splunk_ao.evaluator.validate_code_scorer_scorers_code_validate_post") + @patch("splunk_ao.evaluator.SplunkAOConfig") def test_custom_backoff_multiplier_is_applied( self, mock_config, @@ -1975,7 +1975,7 @@ def test_custom_backoff_multiplier_is_applied( # Time progression mock_time.side_effect = [0.0, 1.0, 2.0, 3.0, 4.0] - metric = CodeMetric(name="Test Backoff Multiplier Metric", node_level=StepType.llm).load_code(str(code_file)) + metric = CodeEvaluator(name="Test Backoff Multiplier Evaluator", node_level=StepType.llm).load_code(str(code_file)) try: metric.create() diff --git a/tests/test_metric_types.py b/tests/test_evaluator_types.py similarity index 60% rename from tests/test_metric_types.py rename to tests/test_evaluator_types.py index 279fc544..ea3bebec 100644 --- a/tests/test_metric_types.py +++ b/tests/test_evaluator_types.py @@ -2,10 +2,10 @@ Tests for the refactored metric type hierarchy. This module tests the four metric types: -- LlmMetric: Custom LLM-based metrics with prompt templates -- LocalMetric: Local function-based metrics -- CodeMetric: Code-based metrics (limited support) -- SplunkAOMetric: Built-in Galileo scorers +- LlmEvaluator: Custom LLM-based metrics with prompt templates +- LocalEvaluator: Local function-based metrics +- CodeEvaluator: Code-based metrics (limited support) +- SplunkAOEvaluator: Built-in Galileo scorers """ from __future__ import annotations @@ -13,17 +13,17 @@ import pytest from galileo_core.schemas.logging.step import StepType -from splunk_ao.metric import CodeMetric, LlmMetric, LocalMetric, Metric, SplunkAOMetric +from splunk_ao.evaluator import CodeEvaluator, LlmEvaluator, LocalEvaluator, Evaluator, SplunkAOEvaluator from splunk_ao.resources.models import OutputTypeEnum, ScorerTypes from splunk_ao.shared.exceptions import ValidationError class TestLlmMetric: - """Tests for LlmMetric class.""" + """Tests for LlmEvaluator class.""" def test_llm_metric_initialization(self): - """Test basic LlmMetric initialization.""" - metric = LlmMetric( + """Test basic LlmEvaluator initialization.""" + metric = LlmEvaluator( name="test_llm", prompt="Rate this response", model="gpt-4o-mini", @@ -39,24 +39,24 @@ def test_llm_metric_initialization(self): assert metric.description == "Test LLM metric" assert metric.tags == ["test", "quality"] assert metric.scorer_type == ScorerTypes.LLM - assert isinstance(metric, LlmMetric) - assert isinstance(metric, Metric) + assert isinstance(metric, LlmEvaluator) + assert isinstance(metric, Evaluator) def test_llm_metric_with_output_type_string(self): - """Test LlmMetric with string output_type.""" - metric = LlmMetric(name="test_llm", prompt="Rate this", output_type="percentage") + """Test LlmEvaluator with string output_type.""" + metric = LlmEvaluator(name="test_llm", prompt="Rate this", output_type="percentage") assert metric.output_type == OutputTypeEnum.PERCENTAGE def test_llm_metric_with_output_type_enum(self): - """Test LlmMetric with enum output_type.""" - metric = LlmMetric(name="test_llm", prompt="Rate this", output_type=OutputTypeEnum.BOOLEAN) + """Test LlmEvaluator with enum output_type.""" + metric = LlmEvaluator(name="test_llm", prompt="Rate this", output_type=OutputTypeEnum.BOOLEAN) assert metric.output_type == OutputTypeEnum.BOOLEAN def test_llm_metric_backward_compatibility_aliases(self): - """Test LlmMetric with deprecated parameter names.""" - metric = LlmMetric(name="test_llm", user_prompt="Old prompt param", model_name="gpt-3.5-turbo", num_judges=2) + """Test LlmEvaluator with deprecated parameter names.""" + metric = LlmEvaluator(name="test_llm", user_prompt="Old prompt param", model_name="gpt-3.5-turbo", num_judges=2) assert metric.prompt == "Old prompt param" assert metric.model == "gpt-3.5-turbo" @@ -64,7 +64,7 @@ def test_llm_metric_backward_compatibility_aliases(self): def test_llm_metric_new_params_override_old(self): """Test that new parameter names override deprecated ones.""" - metric = LlmMetric( + metric = LlmEvaluator( name="test_llm", prompt="New prompt", user_prompt="Old prompt", @@ -79,51 +79,51 @@ def test_llm_metric_new_params_override_old(self): assert metric.judges == 5 def test_llm_metric_requires_prompt(self): - """Test that LlmMetric requires a prompt.""" + """Test that LlmEvaluator requires a prompt.""" with pytest.raises(ValidationError, match="'prompt'.*must be provided"): - LlmMetric(name="test_llm") + LlmEvaluator(name="test_llm") def test_llm_metric_defaults(self): - """Test LlmMetric default values.""" - metric = LlmMetric(name="test_llm", prompt="Rate this") + """Test LlmEvaluator default values.""" + metric = LlmEvaluator(name="test_llm", prompt="Rate this") assert metric.node_level == StepType.llm assert metric.cot_enabled is True assert metric.output_type == OutputTypeEnum.BOOLEAN def test_llm_metric_custom_node_level(self): - """Test LlmMetric with custom node_level.""" - metric = LlmMetric(name="test_llm", prompt="Rate this", node_level=StepType.workflow) + """Test LlmEvaluator with custom node_level.""" + metric = LlmEvaluator(name="test_llm", prompt="Rate this", node_level=StepType.workflow) assert metric.node_level == StepType.workflow def test_llm_metric_cot_disabled(self): - """Test LlmMetric with chain-of-thought disabled.""" - metric = LlmMetric(name="test_llm", prompt="Rate this", cot_enabled=False) + """Test LlmEvaluator with chain-of-thought disabled.""" + metric = LlmEvaluator(name="test_llm", prompt="Rate this", cot_enabled=False) assert metric.cot_enabled is False def test_llm_metric_repr(self): - """Test LlmMetric string representation.""" - metric = LlmMetric(name="test_llm", prompt="Rate this", model="gpt-4o-mini", judges=3) + """Test LlmEvaluator string representation.""" + metric = LlmEvaluator(name="test_llm", prompt="Rate this", model="gpt-4o-mini", judges=3) repr_str = repr(metric) - assert "LlmMetric" in repr_str + assert "LlmEvaluator" in repr_str assert "test_llm" in repr_str assert "gpt-4o-mini" in repr_str assert "3" in repr_str class TestLocalMetric: - """Tests for LocalMetric class.""" + """Tests for LocalEvaluator class.""" def test_local_metric_initialization(self): - """Test basic LocalMetric initialization.""" + """Test basic LocalEvaluator initialization.""" def my_scorer(trace_or_span): return 0.5 - metric = LocalMetric( + metric = LocalEvaluator( name="test_local", scorer_fn=my_scorer, scorable_types=[StepType.llm], @@ -139,21 +139,21 @@ def my_scorer(trace_or_span): assert metric.description == "Test local metric" assert metric.tags == ["local", "custom"] assert metric.scorer_type is None - assert isinstance(metric, LocalMetric) - assert isinstance(metric, Metric) + assert isinstance(metric, LocalEvaluator) + assert isinstance(metric, Evaluator) def test_local_metric_requires_scorer_fn(self): - """Test that LocalMetric requires a scorer_fn.""" + """Test that LocalEvaluator requires a scorer_fn.""" with pytest.raises(ValidationError, match="'scorer_fn' must be provided"): - LocalMetric(name="test_local", scorer_fn=None) + LocalEvaluator(name="test_local", scorer_fn=None) def test_local_metric_default_types(self): - """Test LocalMetric default scorable and aggregatable types.""" + """Test LocalEvaluator default scorable and aggregatable types.""" def my_scorer(trace_or_span): return 1.0 - metric = LocalMetric(name="test_local", scorer_fn=my_scorer) + metric = LocalEvaluator(name="test_local", scorer_fn=my_scorer) assert metric.scorable_types == [StepType.llm] assert metric.aggregatable_types == [StepType.trace] @@ -164,7 +164,7 @@ def test_local_metric_to_local_metric_config(self): def my_scorer(trace_or_span): return 0.75 - metric = LocalMetric( + metric = LocalEvaluator( name="test_local", scorer_fn=my_scorer, scorable_types=[StepType.llm, StepType.workflow], @@ -179,75 +179,75 @@ def my_scorer(trace_or_span): assert config.aggregatable_types == [StepType.trace] def test_local_metric_repr(self): - """Test LocalMetric string representation.""" + """Test LocalEvaluator string representation.""" def my_scorer(trace_or_span): return 0.5 - metric = LocalMetric(name="test_local", scorer_fn=my_scorer) + metric = LocalEvaluator(name="test_local", scorer_fn=my_scorer) repr_str = repr(metric) - assert "LocalMetric" in repr_str + assert "LocalEvaluator" in repr_str assert "test_local" in repr_str assert "my_scorer" in repr_str class TestCodeMetric: - """Tests for CodeMetric class.""" + """Tests for CodeEvaluator class.""" def test_code_metric_initialization(self, tmp_path): - """Test basic CodeMetric initialization.""" - metric = CodeMetric(name="test_code", description="Test code metric", tags=["code"]) + """Test basic CodeEvaluator initialization.""" + metric = CodeEvaluator(name="test_code", description="Test code metric", tags=["code"]) assert metric.name == "test_code" assert metric.description == "Test code metric" assert metric.tags == ["code"] assert metric.scorer_type == ScorerTypes.CODE - assert isinstance(metric, CodeMetric) - assert isinstance(metric, Metric) + assert isinstance(metric, CodeEvaluator) + assert isinstance(metric, Evaluator) def test_code_metric_create_not_implemented(self, tmp_path): - """Test that CodeMetric.create() is now implemented.""" + """Test that CodeEvaluator.create() is now implemented.""" # Create a temporary code file code_file = tmp_path / "scorer.py" code_file.write_text("def score(trace): return 1.0") - metric = CodeMetric(name="test_code") + metric = CodeEvaluator(name="test_code") - # CodeMetric.create() is now implemented, so this test should be updated + # CodeEvaluator.create() is now implemented, so this test should be updated # to verify it works or test it separately assert hasattr(metric, "create") assert callable(metric.create) class TestSplunkAOMetric: - """Tests for SplunkAOMetric class.""" + """Tests for SplunkAOEvaluator class.""" def test_galileo_metric_initialization(self): - """Test basic SplunkAOMetric initialization.""" - metric = SplunkAOMetric(name="test_galileo", description="Test Galileo metric", tags=["galileo"]) + """Test basic SplunkAOEvaluator initialization.""" + metric = SplunkAOEvaluator(name="test_galileo", description="Test Galileo metric", tags=["galileo"]) assert metric.name == "test_galileo" assert metric.description == "Test Galileo metric" assert metric.tags == ["galileo"] - assert isinstance(metric, SplunkAOMetric) - assert isinstance(metric, Metric) + assert isinstance(metric, SplunkAOEvaluator) + assert isinstance(metric, Evaluator) class TestMetricBase: - """Tests for base Metric class.""" + """Tests for base Evaluator class.""" def test_metric_has_scorers_attribute(self): - """Test that Metric class has scorers attribute.""" - assert hasattr(Metric, "scorers") + """Test that Evaluator class has scorers attribute.""" + assert hasattr(Evaluator, "scorers") def test_metric_scorers_is_builtin_scorers(self): - """Test that Metric.metrics is a BuiltInMetrics instance and legacy 'scorers' still exists.""" - from splunk_ao.metric import BuiltInMetrics + """Test that Evaluator.metrics is a BuiltInEvaluators instance and legacy 'scorers' still exists.""" + from splunk_ao.evaluator import BuiltInEvaluators - assert isinstance(Metric.metrics, BuiltInMetrics) + assert isinstance(Evaluator.metrics, BuiltInEvaluators) # Legacy alias should still exist and point to the same instance - assert Metric.scorers is Metric.metrics + assert Evaluator.scorers is Evaluator.metrics def test_metric_common_attributes(self, tmp_path): """Test that all metric types have common attributes.""" @@ -256,10 +256,10 @@ def my_scorer(trace_or_span): return 0.5 metrics = [ - LlmMetric(name="llm", prompt="Rate this"), - LocalMetric(name="local", scorer_fn=my_scorer), - CodeMetric(name="code"), - SplunkAOMetric(name="galileo"), + LlmEvaluator(name="llm", prompt="Rate this"), + LocalEvaluator(name="local", scorer_fn=my_scorer), + CodeEvaluator(name="code"), + SplunkAOEvaluator(name="galileo"), ] for metric in metrics: @@ -273,39 +273,39 @@ def my_scorer(trace_or_span): assert hasattr(metric, "version") def test_metric_to_legacy_metric(self): - """Test conversion to legacy Metric format.""" - metric = LlmMetric(name="test_llm", prompt="Rate this", version=1) + """Test conversion to legacy Evaluator format.""" + metric = LlmEvaluator(name="test_llm", prompt="Rate this", version=1) legacy = metric.to_legacy_metric() assert legacy.name == "test_llm" assert legacy.version == 1 def test_metric_str_representation(self): - """Test Metric __str__ method.""" - metric = LlmMetric(name="test_llm", prompt="Rate this") + """Test Evaluator __str__ method.""" + metric = LlmEvaluator(name="test_llm", prompt="Rate this") str_repr = str(metric) - assert "LlmMetric" in str_repr + assert "LlmEvaluator" in str_repr assert "test_llm" in str_repr def test_metric_delete_raises_for_local_metric(self): - """Test that delete raises ValidationError for LocalMetric.""" + """Test that delete raises ValidationError for LocalEvaluator.""" def my_scorer(trace_or_span): return 0.5 - metric = LocalMetric(name="test_local", scorer_fn=my_scorer) + metric = LocalEvaluator(name="test_local", scorer_fn=my_scorer) with pytest.raises(ValidationError, match="Local metrics don't exist on the server"): metric.delete() def test_metric_refresh_raises_for_local_metric(self): - """Test that refresh raises ValidationError for LocalMetric.""" + """Test that refresh raises ValidationError for LocalEvaluator.""" def my_scorer(trace_or_span): return 0.5 - metric = LocalMetric(name="test_local", scorer_fn=my_scorer) + metric = LocalEvaluator(name="test_local", scorer_fn=my_scorer) with pytest.raises(ValidationError, match="Local metrics don't exist on the server"): metric.refresh() @@ -315,19 +315,19 @@ class TestMetricInheritance: """Tests for metric type inheritance.""" def test_all_metrics_inherit_from_base(self, tmp_path): - """Test that all metric types inherit from Metric.""" + """Test that all metric types inherit from Evaluator.""" def my_scorer(trace_or_span): return 0.5 - # Create a temporary code file for CodeMetric + # Create a temporary code file for CodeEvaluator code_file = tmp_path / "scorer.py" code_file.write_text("def score(trace): return 1.0") - assert isinstance(LlmMetric(name="llm", prompt="Rate"), Metric) - assert isinstance(LocalMetric(name="local", scorer_fn=my_scorer), Metric) - assert isinstance(CodeMetric(name="code"), Metric) - assert isinstance(SplunkAOMetric(name="galileo"), Metric) + assert isinstance(LlmEvaluator(name="llm", prompt="Rate"), Evaluator) + assert isinstance(LocalEvaluator(name="local", scorer_fn=my_scorer), Evaluator) + assert isinstance(CodeEvaluator(name="code"), Evaluator) + assert isinstance(SplunkAOEvaluator(name="galileo"), Evaluator) def test_metric_type_checking(self, tmp_path): """Test isinstance checks for different metric types.""" @@ -335,34 +335,34 @@ def test_metric_type_checking(self, tmp_path): def my_scorer(trace_or_span): return 0.5 - llm = LlmMetric(name="llm", prompt="Rate") - local = LocalMetric(name="local", scorer_fn=my_scorer) - code = CodeMetric(name="code") - galileo = SplunkAOMetric(name="galileo") + llm = LlmEvaluator(name="llm", prompt="Rate") + local = LocalEvaluator(name="local", scorer_fn=my_scorer) + code = CodeEvaluator(name="code") + galileo = SplunkAOEvaluator(name="galileo") - assert isinstance(llm, LlmMetric) and not isinstance(llm, LocalMetric) - assert isinstance(local, LocalMetric) and not isinstance(local, LlmMetric) - assert isinstance(code, CodeMetric) and not isinstance(code, LlmMetric) - assert isinstance(galileo, SplunkAOMetric) and not isinstance(galileo, LlmMetric) + assert isinstance(llm, LlmEvaluator) and not isinstance(llm, LocalEvaluator) + assert isinstance(local, LocalEvaluator) and not isinstance(local, LlmEvaluator) + assert isinstance(code, CodeEvaluator) and not isinstance(code, LlmEvaluator) + assert isinstance(galileo, SplunkAOEvaluator) and not isinstance(galileo, LlmEvaluator) class TestMetricEdgeCases: """Tests for edge cases and error conditions.""" def test_llm_metric_with_minimal_prompt(self): - """Test LlmMetric with minimal prompt.""" + """Test LlmEvaluator with minimal prompt.""" # Minimal prompt that's not empty - metric = LlmMetric(name="test", prompt="Rate") + metric = LlmEvaluator(name="test", prompt="Rate") assert metric.prompt == "Rate" def test_metric_with_empty_tags(self): """Test metrics with empty tags list.""" - metric = LlmMetric(name="test", prompt="Rate", tags=[]) + metric = LlmEvaluator(name="test", prompt="Rate", tags=[]) assert metric.tags == [] def test_metric_with_none_tags(self): """Test metrics with None tags (should default to empty list).""" - metric = LlmMetric(name="test", prompt="Rate", tags=None) + metric = LlmEvaluator(name="test", prompt="Rate", tags=None) assert metric.tags == [] def test_output_type_mapping(self): @@ -378,21 +378,21 @@ def test_output_type_mapping(self): } for string_type, enum_type in output_types.items(): - metric = LlmMetric(name="test", prompt="Rate", output_type=string_type) + metric = LlmEvaluator(name="test", prompt="Rate", output_type=string_type) assert metric.output_type == enum_type def test_output_type_unknown_string(self): """Test that unknown output_type string defaults to PERCENTAGE.""" - metric = LlmMetric(name="test", prompt="Rate", output_type="unknown_type") + metric = LlmEvaluator(name="test", prompt="Rate", output_type="unknown_type") assert metric.output_type == OutputTypeEnum.PERCENTAGE def test_local_metric_with_multiple_scorable_types(self): - """Test LocalMetric with multiple scorable types.""" + """Test LocalEvaluator with multiple scorable types.""" def my_scorer(trace_or_span): return 0.5 - metric = LocalMetric( + metric = LocalEvaluator( name="test", scorer_fn=my_scorer, scorable_types=[StepType.llm, StepType.workflow, StepType.agent] ) diff --git a/tests/test_metrics.py b/tests/test_evaluators.py similarity index 79% rename from tests/test_metrics.py rename to tests/test_evaluators.py index ac2d9360..b46a95a1 100644 --- a/tests/test_metrics.py +++ b/tests/test_evaluators.py @@ -6,7 +6,7 @@ import pytest from galileo_core.schemas.logging.step import StepType -from splunk_ao.metrics import Metrics, create_custom_llm_metric, delete_metric, get_metrics +from splunk_ao.evaluators import Evaluators, create_custom_llm_evaluator, delete_evaluator, get_evaluators from splunk_ao.resources.models import ( BucketedMetrics, HTTPValidationError, @@ -78,10 +78,10 @@ def _log_records_metrics_response_factory() -> LogRecordsMetricsResponse: class TestMetrics: - """Test cases for the Metrics class.""" + """Test cases for the Evaluators class.""" - @patch("splunk_ao.metrics.create_llm_scorer_version_scorers_scorer_id_version_llm_post") - @patch("splunk_ao.metrics.create_scorers_post") + @patch("splunk_ao.evaluators.create_llm_scorer_version_scorers_scorer_id_version_llm_post") + @patch("splunk_ao.evaluators.create_scorers_post") def test_create_custom_llm_metric_success( self, mock_create_scorer, mock_create_version, mock_scorer_response, mock_scorer_version_response ) -> None: @@ -90,10 +90,10 @@ def test_create_custom_llm_metric_success( mock_create_scorer.sync.return_value = mock_scorer_response mock_create_version.sync.return_value = mock_scorer_version_response - metrics = Metrics() + metrics = Evaluators() # Test with default parameters - result = metrics.create_custom_llm_metric(name="test_metric", user_prompt="Rate the quality of this response") + result = metrics.create_custom_llm_evaluator(name="test_metric", user_prompt="Rate the quality of this response") # Verify the result assert result == mock_scorer_version_response @@ -124,8 +124,8 @@ def test_create_custom_llm_metric_success( assert version_request.user_prompt == "Rate the quality of this response" assert create_version_call.kwargs["scorer_id"] == mock_scorer_response.id - @patch("splunk_ao.metrics.create_llm_scorer_version_scorers_scorer_id_version_llm_post") - @patch("splunk_ao.metrics.create_scorers_post") + @patch("splunk_ao.evaluators.create_llm_scorer_version_scorers_scorer_id_version_llm_post") + @patch("splunk_ao.evaluators.create_scorers_post") def test_create_custom_llm_metric_with_custom_parameters( self, mock_create_scorer, mock_create_version, mock_scorer_response, mock_scorer_version_response ) -> None: @@ -134,10 +134,10 @@ def test_create_custom_llm_metric_with_custom_parameters( mock_create_scorer.sync.return_value = mock_scorer_response mock_create_version.sync.return_value = mock_scorer_version_response - metrics = Metrics() + metrics = Evaluators() # Test with custom parameters - result = metrics.create_custom_llm_metric( + result = metrics.create_custom_llm_evaluator( name="custom_metric", user_prompt="Custom prompt for evaluation", node_level=StepType.workflow, @@ -169,24 +169,24 @@ def test_create_custom_llm_metric_with_custom_parameters( version_request = mock_create_version.sync.call_args.kwargs["body"] assert version_request.user_prompt == "Custom prompt for evaluation" - @patch("splunk_ao.metrics.create_llm_scorer_version_scorers_scorer_id_version_llm_post") - @patch("splunk_ao.metrics.create_scorers_post") + @patch("splunk_ao.evaluators.create_llm_scorer_version_scorers_scorer_id_version_llm_post") + @patch("splunk_ao.evaluators.create_scorers_post") def test_create_custom_llm_metric_scorer_creation_failure(self, mock_create_scorer, mock_create_version) -> None: """Test handling of scorer creation failure.""" # Setup mock to raise exception mock_create_scorer.sync.side_effect = Exception("Scorer creation failed") - metrics = Metrics() + metrics = Evaluators() # Test that exception is propagated with pytest.raises(Exception, match="Scorer creation failed"): - metrics.create_custom_llm_metric(name="test_metric", user_prompt="Test prompt") + metrics.create_custom_llm_evaluator(name="test_metric", user_prompt="Test prompt") # Verify create_version was not called mock_create_version.sync.assert_not_called() - @patch("splunk_ao.metrics.create_llm_scorer_version_scorers_scorer_id_version_llm_post") - @patch("splunk_ao.metrics.create_scorers_post") + @patch("splunk_ao.evaluators.create_llm_scorer_version_scorers_scorer_id_version_llm_post") + @patch("splunk_ao.evaluators.create_scorers_post") def test_create_custom_llm_metric_version_creation_failure( self, mock_create_scorer, mock_create_version, mock_scorer_response ) -> None: @@ -195,19 +195,19 @@ def test_create_custom_llm_metric_version_creation_failure( mock_create_scorer.sync.return_value = mock_scorer_response mock_create_version.sync.side_effect = Exception("Version creation failed") - metrics = Metrics() + metrics = Evaluators() # Test that exception is propagated with pytest.raises(Exception, match="Version creation failed"): - metrics.create_custom_llm_metric(name="test_metric", user_prompt="Test prompt") + metrics.create_custom_llm_evaluator(name="test_metric", user_prompt="Test prompt") # Verify create_scorer was called but create_version failed mock_create_scorer.sync.assert_called_once() mock_create_version.sync.assert_called_once() - @patch("splunk_ao.metrics.create_llm_scorer_version_scorers_scorer_id_version_llm_post") - @patch("splunk_ao.metrics.create_scorers_post") - @patch("splunk_ao.metrics._logger") + @patch("splunk_ao.evaluators.create_llm_scorer_version_scorers_scorer_id_version_llm_post") + @patch("splunk_ao.evaluators.create_scorers_post") + @patch("splunk_ao.evaluators._logger") def test_create_custom_llm_metric_logging( self, mock_logger, mock_create_scorer, mock_create_version, mock_scorer_response, mock_scorer_version_response ) -> None: @@ -216,21 +216,21 @@ def test_create_custom_llm_metric_logging( mock_create_scorer.sync.return_value = mock_scorer_response mock_create_version.sync.return_value = mock_scorer_version_response - metrics = Metrics() + metrics = Evaluators() # Create metric - metrics.create_custom_llm_metric(name="test_metric", user_prompt="Test prompt") + metrics.create_custom_llm_evaluator(name="test_metric", user_prompt="Test prompt") # Verify logging was called mock_logger.info.assert_called_once_with("Created custom LLM metric: %s", "test_metric") - @patch("splunk_ao.metrics.delete_scorer_scorers_scorer_id_delete") + @patch("splunk_ao.evaluators.delete_scorer_scorers_scorer_id_delete") @patch("splunk_ao.scorers.Scorers.list") def test_delete_metric_success(self, mock_list_scorers, mock_delete_scorer, mock_scorer_response) -> None: """Test successful deletion of a metric.""" mock_list_scorers.return_value = [mock_scorer_response] - metrics = Metrics() - metrics.delete_metric(name="test_metric") + metrics = Evaluators() + metrics.delete_evaluator(name="test_metric") mock_list_scorers.assert_called_once_with(name="test_metric") mock_delete_scorer.sync.assert_called_once_with( @@ -241,37 +241,37 @@ def test_delete_metric_success(self, mock_list_scorers, mock_delete_scorer, mock def test_delete_metric_not_found(self, mock_list_scorers) -> None: """Test deleting a metric that does not exist.""" mock_list_scorers.return_value = [] - metrics = Metrics() + metrics = Evaluators() with pytest.raises(ValueError, match="Scorer with name test_metric not found."): - metrics.delete_metric(name="test_metric") + metrics.delete_evaluator(name="test_metric") - @patch("splunk_ao.metrics.delete_scorer_scorers_scorer_id_delete") + @patch("splunk_ao.evaluators.delete_scorer_scorers_scorer_id_delete") @patch("splunk_ao.scorers.Scorers.list") def test_delete_metric_api_failure(self, mock_list_scorers, mock_delete_scorer, mock_scorer_response) -> None: """Test API failure when deleting a metric.""" mock_list_scorers.return_value = [mock_scorer_response] mock_delete_scorer.sync.return_value = None - metrics = Metrics() + metrics = Evaluators() with pytest.raises(ValueError, match="Failed to delete metric."): - metrics.delete_metric(name="test_metric") + metrics.delete_evaluator(name="test_metric") class TestPublicFunctions: """Test cases for public functions.""" - @patch("splunk_ao.metrics.Metrics") + @patch("splunk_ao.evaluators.Evaluators") def test_create_custom_llm_metric_function(self, mock_metrics_class) -> None: - """Test the public create_custom_llm_metric function.""" + """Test the public create_custom_llm_evaluator function.""" # Setup mock mock_metrics_instance = Mock() mock_metrics_class.return_value = mock_metrics_instance mock_result = Mock(spec=BaseScorerVersionResponse) - mock_metrics_instance.create_custom_llm_metric.return_value = mock_result + mock_metrics_instance.create_custom_llm_evaluator.return_value = mock_result # Call the public function - result = create_custom_llm_metric( + result = create_custom_llm_evaluator( name="test_metric", user_prompt="Test prompt", node_level=StepType.workflow, @@ -284,11 +284,11 @@ def test_create_custom_llm_metric_function(self, mock_metrics_class) -> None: ground_truth=True, ) - # Verify Metrics class was instantiated + # Verify Evaluators class was instantiated mock_metrics_class.assert_called_once() # Verify the method was called with correct parameters - mock_metrics_instance.create_custom_llm_metric.assert_called_once_with( + mock_metrics_instance.create_custom_llm_evaluator.assert_called_once_with( "test_metric", "Test prompt", StepType.workflow, @@ -304,20 +304,20 @@ def test_create_custom_llm_metric_function(self, mock_metrics_class) -> None: # Verify the result is returned assert result == mock_result - @patch("splunk_ao.metrics.Metrics") + @patch("splunk_ao.evaluators.Evaluators") def test_create_custom_llm_metric_function_default_parameters(self, mock_metrics_class) -> None: """Test the public function with default parameters.""" # Setup mock mock_metrics_instance = Mock() mock_metrics_class.return_value = mock_metrics_instance mock_result = Mock(spec=BaseScorerVersionResponse) - mock_metrics_instance.create_custom_llm_metric.return_value = mock_result + mock_metrics_instance.create_custom_llm_evaluator.return_value = mock_result # Call the public function with minimal parameters - result = create_custom_llm_metric(name="test_metric", user_prompt="Test prompt") + result = create_custom_llm_evaluator(name="test_metric", user_prompt="Test prompt") # Verify the method was called with default parameters - mock_metrics_instance.create_custom_llm_metric.assert_called_once_with( + mock_metrics_instance.create_custom_llm_evaluator.assert_called_once_with( "test_metric", "Test prompt", StepType.llm, # default @@ -333,22 +333,22 @@ def test_create_custom_llm_metric_function_default_parameters(self, mock_metrics # Verify the result is returned assert result == mock_result - @patch("splunk_ao.metrics.Metrics") + @patch("splunk_ao.evaluators.Evaluators") def test_delete_metric_function(self, mock_metrics_class) -> None: - """Test the public delete_metric function.""" + """Test the public delete_evaluator function.""" mock_metrics_instance = Mock() mock_metrics_class.return_value = mock_metrics_instance - delete_metric(name="test_metric") + delete_evaluator(name="test_metric") mock_metrics_class.assert_called_once() - mock_metrics_instance.delete_metric.assert_called_once_with("test_metric") + mock_metrics_instance.delete_evaluator.assert_called_once_with("test_metric") class TestEdgeCases: """Test edge cases and boundary conditions.""" - @patch("splunk_ao.metrics.create_llm_scorer_version_scorers_scorer_id_version_llm_post") - @patch("splunk_ao.metrics.create_scorers_post") + @patch("splunk_ao.evaluators.create_llm_scorer_version_scorers_scorer_id_version_llm_post") + @patch("splunk_ao.evaluators.create_scorers_post") def test_empty_string_parameters( self, mock_create_scorer, mock_create_version, mock_scorer_response, mock_scorer_version_response ) -> None: @@ -357,10 +357,10 @@ def test_empty_string_parameters( mock_create_scorer.sync.return_value = mock_scorer_response mock_create_version.sync.return_value = mock_scorer_version_response - metrics = Metrics() + metrics = Evaluators() # Test with empty strings - result = metrics.create_custom_llm_metric( + result = metrics.create_custom_llm_evaluator( name="", # Empty name user_prompt="", # Empty prompt description="", @@ -378,8 +378,8 @@ def test_empty_string_parameters( version_request = mock_create_version.sync.call_args.kwargs["body"] assert version_request.user_prompt == "" - @patch("splunk_ao.metrics.create_llm_scorer_version_scorers_scorer_id_version_llm_post") - @patch("splunk_ao.metrics.create_scorers_post") + @patch("splunk_ao.evaluators.create_llm_scorer_version_scorers_scorer_id_version_llm_post") + @patch("splunk_ao.evaluators.create_scorers_post") def test_large_num_judges( self, mock_create_scorer, mock_create_version, mock_scorer_response, mock_scorer_version_response ) -> None: @@ -388,10 +388,10 @@ def test_large_num_judges( mock_create_scorer.sync.return_value = mock_scorer_response mock_create_version.sync.return_value = mock_scorer_version_response - metrics = Metrics() + metrics = Evaluators() # Test with large number of judges - result = metrics.create_custom_llm_metric(name="test_metric", user_prompt="Test prompt", num_judges=100) + result = metrics.create_custom_llm_evaluator(name="test_metric", user_prompt="Test prompt", num_judges=100) # Verify the result assert result == mock_scorer_version_response @@ -403,8 +403,8 @@ def test_large_num_judges( version_request = mock_create_version.sync.call_args.kwargs["body"] assert version_request.user_prompt == "Test prompt" - @patch("splunk_ao.metrics.create_llm_scorer_version_scorers_scorer_id_version_llm_post") - @patch("splunk_ao.metrics.create_scorers_post") + @patch("splunk_ao.evaluators.create_llm_scorer_version_scorers_scorer_id_version_llm_post") + @patch("splunk_ao.evaluators.create_scorers_post") def test_long_tag_list( self, mock_create_scorer, mock_create_version, mock_scorer_response, mock_scorer_version_response ) -> None: @@ -413,11 +413,11 @@ def test_long_tag_list( mock_create_scorer.sync.return_value = mock_scorer_response mock_create_version.sync.return_value = mock_scorer_version_response - metrics = Metrics() + metrics = Evaluators() # Test with many tags long_tag_list = [f"tag_{i}" for i in range(50)] - result = metrics.create_custom_llm_metric(name="test_metric", user_prompt="Test prompt", tags=long_tag_list) + result = metrics.create_custom_llm_evaluator(name="test_metric", user_prompt="Test prompt", tags=long_tag_list) # Verify the result assert result == mock_scorer_version_response @@ -429,7 +429,7 @@ def test_long_tag_list( class TestGetMetrics: - @patch("splunk_ao.metrics.query_metrics_projects_project_id_metrics_search_post.sync") + @patch("splunk_ao.evaluators.query_metrics_projects_project_id_metrics_search_post.sync") def test_successful_call(self, mock_api_call): mock_response = _log_records_metrics_response_factory() mock_api_call.return_value = mock_response @@ -437,13 +437,13 @@ def test_successful_call(self, mock_api_call): start_time = datetime.datetime.now() end_time = start_time + datetime.timedelta(hours=1) - response = get_metrics(project_id=FIXED_PROJECT_ID, start_time=start_time, end_time=end_time) + response = get_evaluators(project_id=FIXED_PROJECT_ID, start_time=start_time, end_time=end_time) mock_api_call.assert_called_once() assert FIXED_PROJECT_ID in mock_api_call.call_args[1]["project_id"] assert response == mock_response - @patch("splunk_ao.metrics.query_metrics_projects_project_id_metrics_search_post.sync") + @patch("splunk_ao.evaluators.query_metrics_projects_project_id_metrics_search_post.sync") def test_api_failure_raises_value_error(self, mock_api_call): mock_api_call.return_value = None @@ -451,11 +451,11 @@ def test_api_failure_raises_value_error(self, mock_api_call): end_time = start_time + datetime.timedelta(hours=1) with pytest.raises(ValueError, match="Failed to query for metrics."): - get_metrics(project_id=FIXED_PROJECT_ID, start_time=start_time, end_time=end_time) + get_evaluators(project_id=FIXED_PROJECT_ID, start_time=start_time, end_time=end_time) mock_api_call.assert_called_once() - @patch("splunk_ao.metrics.query_metrics_projects_project_id_metrics_search_post.sync") + @patch("splunk_ao.evaluators.query_metrics_projects_project_id_metrics_search_post.sync") def test_http_validation_error_raises_exception(self, mock_api_call): detail = [ValidationError(loc=["body", "project_id"], msg="value is not a valid uuid", type_="type_error.uuid")] mock_api_call.return_value = HTTPValidationError(detail=detail) @@ -464,9 +464,9 @@ def test_http_validation_error_raises_exception(self, mock_api_call): end_time = start_time + datetime.timedelta(hours=1) with pytest.raises(ValueError, match=re.escape(str(detail))): - get_metrics(project_id=FIXED_PROJECT_ID, start_time=start_time, end_time=end_time) + get_evaluators(project_id=FIXED_PROJECT_ID, start_time=start_time, end_time=end_time) - @patch("splunk_ao.metrics.query_metrics_projects_project_id_metrics_search_post.sync") + @patch("splunk_ao.evaluators.query_metrics_projects_project_id_metrics_search_post.sync") def test_passes_all_parameters_correctly(self, mock_api_call): mock_api_call.return_value = _log_records_metrics_response_factory() @@ -478,7 +478,7 @@ def test_passes_all_parameters_correctly(self, mock_api_call): group_by = "some_column" interval = 10 - get_metrics( + get_evaluators( project_id=FIXED_PROJECT_ID, start_time=start_time, end_time=end_time, diff --git a/tests/test_experiments.py b/tests/test_experiments.py index 05729f81..b69aa109 100644 --- a/tests/test_experiments.py +++ b/tests/test_experiments.py @@ -793,7 +793,7 @@ def test_run_experiment_no_prompt_no_dataset_raises( with pytest.raises(ValueError, match="dataset"): run_experiment("test_experiment", project="awesome-new-project") - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.object(splunk_ao.datasets.Datasets, "get") @@ -1162,7 +1162,7 @@ def test_experiments_run_with_prompt_settings_as_dict(self, mock_create: Mock) - assert ps.max_tokens == 256 @travel(datetime(2012, 1, 1), tick=False) - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.object(splunk_ao.datasets.Datasets, "get") @@ -1289,7 +1289,7 @@ def test_run_experiment_on_error_warns_when_unused_in_prompt_template_flow( ) on_error.assert_not_called() - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.object(splunk_ao.datasets.Datasets, "get") @@ -1554,7 +1554,7 @@ def test_run_experiment_with_experiment_tags_basic( assert mock_upsert_tag.call_count == 3 - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.object(splunk_ao.datasets.Datasets, "get") @@ -1601,7 +1601,7 @@ def test_run_experiment_with_dataset_limit( total_traces += len(payload.traces) assert total_traces == 150 - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.object(splunk_ao.experiments.Experiments, "create", return_value=experiment_response()) @@ -1662,7 +1662,7 @@ def simple_function(input: dict) -> str: assert '{"input": "Which continent is Spain in?"}' in span_inputs[0] assert '{"input": "Which continent is Japan in?"}' in span_inputs[1] - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") @patch.object(splunk_ao.datasets.Datasets, "get") diff --git a/tests/test_export.py b/tests/test_export.py index 96780221..7ac0b603 100644 --- a/tests/test_export.py +++ b/tests/test_export.py @@ -6,7 +6,7 @@ import pytest from splunk_ao.export import export_records -from splunk_ao.log_streams import LogStream +from splunk_ao.agent_streams import AgentStream from splunk_ao.resources.errors import UnexpectedStatus from splunk_ao.resources.models import ( LLMExportFormat, @@ -64,7 +64,7 @@ def test_export_records_with_defaults(mock_export_records_stream): assert request_body.sort == LogRecordsSortClause(column_id="created_at", ascending=False) -@patch("splunk_ao.export.LogStreams._list_all") +@patch("splunk_ao.export.AgentStreams._list_all") @patch("splunk_ao.export.export_records_stream") def test_export_records_default_log_stream(mock_export_records_stream, mock_log_streams_list_all): project_id = str(uuid4()) @@ -72,15 +72,15 @@ def test_export_records_default_log_stream(mock_export_records_stream, mock_log_ now = datetime.now() # Create mock log streams, ensuring one is clearly the oldest - log_stream_1 = LogStream() + log_stream_1 = AgentStream() log_stream_1.id = str(uuid4()) log_stream_1.created_at = now - log_stream_2 = LogStream() + log_stream_2 = AgentStream() log_stream_2.id = oldest_log_stream_id log_stream_2.created_at = now - timedelta(days=1) - log_stream_3 = LogStream() + log_stream_3 = AgentStream() log_stream_3.id = str(uuid4()) log_stream_3.created_at = now + timedelta(days=1) @@ -99,7 +99,7 @@ def test_export_records_default_log_stream(mock_export_records_stream, mock_log_ assert request_body.experiment_id is None -@patch("splunk_ao.export.LogStreams._list_all") +@patch("splunk_ao.export.AgentStreams._list_all") @patch("splunk_ao.export.export_records_stream") def test_export_records_no_default_log_stream(mock_export_records_stream, mock_log_streams_list_all): project_id = str(uuid4()) diff --git a/tests/test_galileo_context.py b/tests/test_galileo_context.py index 43fd3e39..3eb96be8 100644 --- a/tests/test_galileo_context.py +++ b/tests/test_galileo_context.py @@ -12,7 +12,7 @@ def reset_context() -> None: splunk_ao_context.reset() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_nested_context_restoration( @@ -65,7 +65,7 @@ def test_nested_context_restoration( assert _experiment_id_context.get() is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_context_update_with_defaults( diff --git a/tests/test_langchain.py b/tests/test_langchain.py index d97343f2..0276852b 100644 --- a/tests/test_langchain.py +++ b/tests/test_langchain.py @@ -24,7 +24,7 @@ class TestSplunkAOCallback: @pytest.fixture - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def splunk_ao_logger(self, mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock): @@ -1023,7 +1023,7 @@ class TestSplunkAOCallbackWithIngestionHook: @pytest.fixture(autouse=True) def logger_mocks(self): with ( - patch("splunk_ao.logger.logger.LogStreams") as mock_logstreams, + patch("splunk_ao.logger.logger.AgentStreams") as mock_logstreams, patch("splunk_ao.logger.logger.Projects") as mock_projects, patch("splunk_ao.logger.logger.Traces") as mock_traces, ): diff --git a/tests/test_langchain_async.py b/tests/test_langchain_async.py index ce69f066..8e19b955 100644 --- a/tests/test_langchain_async.py +++ b/tests/test_langchain_async.py @@ -21,7 +21,7 @@ class TestSplunkAOAsyncCallback: @pytest.fixture - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def splunk_ao_logger(self, mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock): diff --git a/tests/test_langchain_middleware.py b/tests/test_langchain_middleware.py index fa68996f..b5128eec 100644 --- a/tests/test_langchain_middleware.py +++ b/tests/test_langchain_middleware.py @@ -88,7 +88,7 @@ class RealArgsSchema(BaseModel): @pytest.fixture -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def splunk_ao_logger(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock): @@ -494,7 +494,7 @@ class TestIngestionHook: @pytest.fixture(autouse=True) def logger_mocks(self): with ( - patch("splunk_ao.logger.logger.LogStreams") as mock_logstreams, + patch("splunk_ao.logger.logger.AgentStreams") as mock_logstreams, patch("splunk_ao.logger.logger.Projects") as mock_projects, patch("splunk_ao.logger.logger.Traces") as mock_traces, ): diff --git a/tests/test_logger_batch.py b/tests/test_logger_batch.py index 87f624e5..b42516c5 100644 --- a/tests/test_logger_batch.py +++ b/tests/test_logger_batch.py @@ -78,7 +78,7 @@ def test_disable_splunk_ao_logger(mock_traces_client: Mock, monkeypatch, caplog, mock_traces_client.ingest_traces.assert_not_called() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_single_span_trace_to_galileo( @@ -125,7 +125,7 @@ def test_single_span_trace_to_galileo( assert logger._parent_stack == deque() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_all_span_types_with_redacted_fields( @@ -288,7 +288,7 @@ def test_single_span_trace_to_galileo_experiment_id( assert logger._parent_stack == deque() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_nested_span_trace_to_galileo( @@ -339,7 +339,7 @@ def test_nested_span_trace_to_galileo( assert logger._parent_stack == deque() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_add_agent_span(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock) -> None: @@ -369,7 +369,7 @@ def test_add_agent_span(mock_traces_client: Mock, mock_projects_client: Mock, mo assert logger._parent_stack == deque() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_multi_span_trace_to_galileo( @@ -437,7 +437,7 @@ def test_multi_span_trace_to_galileo( @pytest.mark.asyncio -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") async def test_single_span_trace_to_galileo_with_async( @@ -494,7 +494,7 @@ def local_scorer(step: Trace | Span) -> int: assert logger._parent_stack == deque() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_retriever_span_str_output( @@ -520,7 +520,7 @@ def test_retriever_span_str_output( assert payload.traces[0].spans[0].output == [Document(content="response", metadata=None)] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_retriever_span_list_str_output( @@ -549,7 +549,7 @@ def test_retriever_span_list_str_output( ] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_retriever_span_dict_output( @@ -584,7 +584,7 @@ def test_retriever_span_dict_output( assert payload.traces[0].spans[1].output == [Document(content="response2", metadata={"key": "value"})] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_retriever_span_list_dict_output( @@ -631,7 +631,7 @@ def test_retriever_span_list_dict_output( ] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_retriever_span_document_output( @@ -661,7 +661,7 @@ def test_retriever_span_document_output( assert payload.traces[0].spans[0].output == [Document(content="response", metadata={"key": "value"})] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_retriever_span_list_document_output( @@ -694,7 +694,7 @@ def test_retriever_span_list_document_output( ] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_retriever_span_none_output( @@ -718,7 +718,7 @@ def test_retriever_span_none_output( assert payload.traces[0].spans[0].output == [Document(content="", metadata={})] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_all_spans(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock) -> None: @@ -757,7 +757,7 @@ def test_conclude_all_spans(mock_traces_client: Mock, mock_projects_client: Mock assert logger._parent_stack == deque() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_flush_with_conclude_all_spans( @@ -804,7 +804,7 @@ def test_flush_with_conclude_all_spans( assert logger._parent_stack == deque() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_flush_workflow_keeps_message_trace_gets_string( @@ -1187,7 +1187,7 @@ def test_get_last_output_llm_message_raw() -> None: assert output.role == MessageRole.assistant -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_session_create(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock) -> None: @@ -1209,7 +1209,7 @@ def test_session_create(mock_traces_client: Mock, mock_projects_client: Mock, mo assert logger.session_id == session_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9c" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_session_create_with_metadata( @@ -1235,7 +1235,7 @@ def test_session_create_with_metadata( assert logger.session_id == session_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9c" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_session_create_empty_values( @@ -1257,7 +1257,7 @@ def test_session_create_empty_values( assert logger.session_id == session_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9c" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_session_clear(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock) -> None: @@ -1277,7 +1277,7 @@ def test_session_clear(mock_traces_client: Mock, mock_projects_client: Mock, moc assert logger.session_id is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_session_id_on_flush( @@ -1303,7 +1303,7 @@ def test_session_id_on_flush( assert str(payload.session_id) == session_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9c" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_set_session_id(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock) -> None: @@ -1329,7 +1329,7 @@ def test_set_session_id(mock_traces_client: Mock, mock_projects_client: Mock, mo assert payload.session_id == UUID(session_id) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_start_session_with_external_id( @@ -1401,7 +1401,7 @@ def test_start_session_with_external_id( assert payload.session_id == session_id -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_logger_init_with_project_id_and_log_stream_id( @@ -1422,7 +1422,7 @@ def test_logger_init_with_project_id_and_log_stream_id( assert logger.log_stream_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9b" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_logger_init_with_project_id_and_log_stream_name( @@ -1441,7 +1441,7 @@ def test_logger_init_with_project_id_and_log_stream_name( assert logger.log_stream_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9b" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_logger_init_with_project_name_and_log_stream_id( @@ -1460,7 +1460,7 @@ def test_logger_init_with_project_name_and_log_stream_id( assert logger.log_stream_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9b" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_logger_init_with_project_name_and_experiment_id( @@ -1480,7 +1480,7 @@ def test_logger_init_with_project_name_and_experiment_id( assert logger.experiment_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9b" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_logger_init_with_project_id_and_experiment_id( @@ -1502,7 +1502,7 @@ def test_logger_init_with_project_id_and_experiment_id( assert logger.experiment_id == "6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9b" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_ingestion_hook_sync( @@ -1527,7 +1527,7 @@ def test_ingestion_hook_sync( @pytest.mark.asyncio -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") async def test_ingestion_hook_async( @@ -1552,7 +1552,7 @@ async def test_ingestion_hook_async( @pytest.mark.asyncio -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") async def test_ingest_traces_methods( @@ -1573,7 +1573,7 @@ async def test_ingest_traces_methods( assert mock_traces_client_instance.ingest_traces.call_count == 2 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_ingestion_hook_with_real_redaction( @@ -1619,7 +1619,7 @@ def redact_and_forward(ingest_request: TracesIngestRequest): assert payload.traces[0].input == "This is a [REDACTED]" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_add_single_llm_span_trace_ingestion( @@ -1668,7 +1668,7 @@ def test_add_single_llm_span_trace_ingestion( assert logger._parent_stack == deque() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_flush_with_unconcluded_trace_redaction( @@ -1745,7 +1745,7 @@ def test_get_last_output_with_redacted_output() -> None: ), ], ) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_start_trace_auto_conversion( @@ -1776,7 +1776,7 @@ def test_start_trace_auto_conversion( assert getattr(payload_trace, attr) == expected_value, f"payload.{attr} mismatch" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_multimodal_input_not_stringified_at_trace_level( @@ -1839,7 +1839,7 @@ def test_multimodal_input_not_stringified_at_trace_level( ), ], ) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_start_trace_valid_input_types( @@ -2001,7 +2001,7 @@ class TestMultipleLoggerInstanceIsolation: on one logger do not affect another logger's state or trace hierarchy. """ - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_loggers_have_isolated_state( @@ -2057,7 +2057,7 @@ def test_loggers_have_isolated_state( assert isinstance(workflow_a, WorkflowSpan) and len(workflow_a.spans) == 1 assert workflow_a.spans[0].name == "llm_a" - @patch("splunk_ao.logger.logger.LogStreams") + @patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_reset_only_affects_own_logger( @@ -2170,7 +2170,7 @@ def record_agent_control_auto_enable(_logger): assert calls == [("atexit", "terminate"), ("agent_control", None)] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_standard_init_registers_atexit_before_agent_control_auto_enable( @@ -2205,7 +2205,7 @@ def record_agent_control_auto_enable(_logger): assert calls == [("atexit", "terminate"), ("agent_control", None)] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_flush_does_not_propagate_exceptions( @@ -2234,7 +2234,7 @@ def test_flush_does_not_propagate_exceptions( assert result is None or result == [] -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_terminate_does_not_propagate_exceptions( @@ -2265,7 +2265,7 @@ def test_terminate_does_not_propagate_exceptions( pytest.fail(f"terminate() should not propagate exceptions, but raised: {e}") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_ingest_traces_lazy_creates_client_for_ingestion_hook( @@ -2306,7 +2306,7 @@ def capture_hook(ingest_request: TracesIngestRequest) -> None: mock_traces_instance.ingest_traces.assert_called_once_with(captured_payload) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_ingest_traces_reuses_existing_client( diff --git a/tests/test_logger_distributed.py b/tests/test_logger_distributed.py index df479e8f..6aa5f164 100644 --- a/tests/test_logger_distributed.py +++ b/tests/test_logger_distributed.py @@ -61,7 +61,7 @@ def test_disable_splunk_ao_logger(mock_traces_client: Mock, monkeypatch, caplog) mock_traces_client.update_span.assert_not_called() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_start_trace(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock) -> None: @@ -96,7 +96,7 @@ def test_start_trace(mock_traces_client: Mock, mock_projects_client: Mock, mock_ @patch("splunk_ao.logger.logger.IngestTraces") @patch("splunk_ao.logger.logger.Traces") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") def test_distributed_logger_uses_ingest_client_when_ingest_service_is_available( mock_projects_client: Mock, mock_logstreams_client: Mock, mock_traces_client: Mock, mock_ingest_traces_client: Mock @@ -115,7 +115,7 @@ def test_distributed_logger_uses_ingest_client_when_ingest_service_is_available( mock_traces_client.assert_not_called() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_nested_distributed_spans_submit_independently( @@ -142,7 +142,7 @@ def test_nested_distributed_spans_submit_independently( assert span_tasks[1].request.parent_id == workflow.id -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_add_llm_span(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock) -> None: @@ -209,7 +209,7 @@ def test_add_llm_span(mock_traces_client: Mock, mock_projects_client: Mock, mock assert request.spans[0].step_number == 1 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace(mock_traces_client: Mock, mock_projects_client: Mock, mock_logstreams_client: Mock) -> None: @@ -263,7 +263,7 @@ def test_conclude_trace(mock_traces_client: Mock, mock_projects_client: Mock, mo assert request.duration_ns == 1_000_000 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace_with_span( @@ -350,7 +350,7 @@ def test_conclude_trace_with_span( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace_and_start_new_trace( @@ -469,7 +469,7 @@ def test_conclude_trace_and_start_new_trace( assert request.traces[0].metrics.duration_ns is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace_with_nested_span( @@ -602,7 +602,7 @@ def test_conclude_trace_with_nested_span( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_all_with_nested_span( @@ -725,7 +725,7 @@ def test_conclude_all_with_nested_span( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace_with_agent_span( @@ -861,7 +861,7 @@ def test_conclude_trace_with_agent_span( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_trace_with_multiple_nested_spans( @@ -1084,7 +1084,7 @@ def test_trace_with_multiple_nested_spans( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_trace_with_nested_span_and_sibling( @@ -1238,7 +1238,7 @@ def test_trace_with_nested_span_and_sibling( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_add_llm_span_and_conclude_existing_trace( @@ -1316,7 +1316,7 @@ def test_add_llm_span_and_conclude_existing_trace( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_add_nested_span_and_conclude_existing_trace( @@ -1440,7 +1440,7 @@ def test_add_nested_span_and_conclude_existing_trace( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_add_llm_span_and_conclude_existing_workflow_span( @@ -1521,7 +1521,7 @@ def test_add_llm_span_and_conclude_existing_workflow_span( assert request.status_code == 200 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_add_nested_span_and_conclude_existing_span( @@ -1648,7 +1648,7 @@ def test_add_nested_span_and_conclude_existing_span( assert request.status_code == 200 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_catch_error_trace_span_ids_in_batch_mode( @@ -1667,7 +1667,7 @@ def test_catch_error_trace_span_ids_in_batch_mode( SplunkAOLogger(project="my_project", log_stream="my_log_stream", span_id="6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9e") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_catch_error_mismatched_trace_span_ids( @@ -1702,7 +1702,7 @@ def test_catch_error_mismatched_trace_span_ids( assert logger._parent_stack[1].id == UUID("6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9e") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_get_tracing_headers_with_workflow_span( @@ -1727,7 +1727,7 @@ def test_get_tracing_headers_with_workflow_span( assert headers["Splunk-AO-Parent-ID"] == str(workflow_span.id) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_get_tracing_headers_with_agent_span( @@ -1752,7 +1752,7 @@ def test_get_tracing_headers_with_agent_span( assert headers["Splunk-AO-Parent-ID"] == str(agent_span.id) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_get_tracing_headers_batch_mode_error( @@ -1772,7 +1772,7 @@ def test_get_tracing_headers_batch_mode_error( assert "only supported in distributed mode" in str(exc_info.value) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_get_tracing_headers_no_trace_error( @@ -1791,7 +1791,7 @@ def test_get_tracing_headers_no_trace_error( assert "Start trace before getting tracing headers" in str(exc_info.value) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_update_trace_output_and_duration_streaming( @@ -1840,7 +1840,7 @@ def test_update_trace_output_and_duration_streaming( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_update_span_output_and_duration_streaming( @@ -1891,7 +1891,7 @@ def test_update_span_output_and_duration_streaming( assert request.status_code == 200 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_update_trace_with_none_duration( @@ -1930,7 +1930,7 @@ def test_update_trace_with_none_duration( assert request.duration_ns is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_update_span_with_none_duration( @@ -1971,7 +1971,7 @@ def test_update_span_with_none_duration( assert request.duration_ns is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_update_trace_and_span_with_duration_in_nested_structure( @@ -2024,7 +2024,7 @@ def test_update_trace_and_span_with_duration_in_nested_structure( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace_inherits_last_llm_child_output( @@ -2075,7 +2075,7 @@ def test_conclude_trace_inherits_last_llm_child_output( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace_with_multiple_llm_children_inherits_last( @@ -2128,7 +2128,7 @@ def test_conclude_trace_with_multiple_llm_children_inherits_last( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace_inherits_last_workflow_span_output( @@ -2185,7 +2185,7 @@ def test_conclude_trace_inherits_last_workflow_span_output( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_trace_explicit_output_overrides_child( @@ -2229,7 +2229,7 @@ def test_conclude_trace_explicit_output_overrides_child( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_conclude_workflow_span_inherits_last_child_output( @@ -2283,7 +2283,7 @@ def test_conclude_workflow_span_inherits_last_child_output( assert request.duration_ns == 3_000_000 -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_distributed_flush_concludes_unconcluded_trace( @@ -2342,7 +2342,7 @@ def test_distributed_flush_concludes_unconcluded_trace( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_distributed_flush_no_op_if_already_concluded( @@ -2398,7 +2398,7 @@ def test_distributed_flush_no_op_if_already_concluded( assert request.is_complete -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_distributed_flush_waits_for_tasks( @@ -2432,7 +2432,7 @@ def test_distributed_flush_waits_for_tasks( assert logger._task_handler.all_tasks_completed() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_terminate_stops_task_handler_when_agent_control_unregister_fails( @@ -2461,7 +2461,7 @@ def test_terminate_stops_task_handler_when_agent_control_unregister_fails( task_handler.terminate.assert_called_once() -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_batch_mode_flush_still_uses_get_last_output( diff --git a/tests/test_logger_timestamps.py b/tests/test_logger_timestamps.py index 56bae4a1..f53ac41d 100644 --- a/tests/test_logger_timestamps.py +++ b/tests/test_logger_timestamps.py @@ -5,7 +5,7 @@ from tests.testutils.setup import setup_mock_logstreams_client, setup_mock_projects_client -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") def test_rapid_span_creation_ensures_uniqueness(mock_projects_client: Mock, mock_logstreams_client: Mock): """Tests that creating spans in a tight loop results in unique, monotonically increasing timestamps.""" @@ -23,7 +23,7 @@ def test_rapid_span_creation_ensures_uniqueness(mock_projects_client: Mock, mock assert timestamps == sorted(timestamps), "Timestamps should be monotonically increasing" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") def test_user_provided_timestamps_are_respected(mock_projects_client: Mock, mock_logstreams_client: Mock): """Tests that timestamps provided by the user are not modified.""" @@ -46,7 +46,7 @@ def test_user_provided_timestamps_are_respected(mock_projects_client: Mock, mock assert span_timestamps == timestamps, "User-provided timestamps should be respected" -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") def test_mixed_default_and_user_timestamps(mock_projects_client: Mock, mock_logstreams_client: Mock): """Tests that the internal state for default timestamp generation is not affected by user-provided timestamps.""" diff --git a/tests/test_middleware_tracing.py b/tests/test_middleware_tracing.py index 28f614ff..81ee7751 100644 --- a/tests/test_middleware_tracing.py +++ b/tests/test_middleware_tracing.py @@ -64,7 +64,7 @@ def client(app): return TestClient(app) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_middleware_extracts_headers( @@ -123,7 +123,7 @@ def test_middleware_extracts_headers( assert data["span_id"] == parent_id -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_middleware_handles_missing_headers( @@ -143,7 +143,7 @@ def test_middleware_handles_missing_headers( assert data["span_id"] is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_middleware_handles_partial_headers( @@ -172,7 +172,7 @@ def test_middleware_handles_partial_headers( assert data["span_id"] is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_context_cleanup_after_request( @@ -201,7 +201,7 @@ def test_context_cleanup_after_request( assert data2["parent_id"] is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_get_request_logger_when_parent_id_equals_trace_id( @@ -233,7 +233,7 @@ def test_get_request_logger_when_parent_id_equals_trace_id( assert data["span_id"] is None -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_mismatched_trace_and_span_ids( @@ -268,7 +268,7 @@ def test_mismatched_trace_and_span_ids( @patch("splunk_ao.logger.logger.Projects") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") def test_invalid_uuid_headers_raise_exception( mock_logstreams_client: Mock, mock_projects_client: Mock, app: FastAPI, client: TestClient ): diff --git a/tests/test_openai.py b/tests/test_openai.py index 4599a8f9..91699e5f 100644 --- a/tests/test_openai.py +++ b/tests/test_openai.py @@ -30,7 +30,7 @@ def openai_incorrect_api_key_error() -> bytes: @patch("openai.resources.chat.Completions.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_basic_openai_call( @@ -94,7 +94,7 @@ def test_basic_openai_call( @patch("openai.resources.chat.Completions.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_streamed_openai_call( @@ -142,7 +142,7 @@ def test_streamed_openai_call( @patch("openai.resources.chat.Completions.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_openai_api_calls_as_parent_span( @@ -195,7 +195,7 @@ def call_openai(model: str = "gpt-3.5-turbo"): "openai.resources.chat.Completions.create", side_effect=openai.OpenAIError("The api_key client option must be set either"), ) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_openai_error_trace( @@ -227,7 +227,7 @@ def call_openai(model: str = "gpt-3.5-turbo"): @patch("openai.resources.chat.Completions.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_openai_error_trace_( @@ -272,7 +272,7 @@ def call_openai(model: str = "gpt-3.5-turbo"): "openai.resources.chat.Completions.create", side_effect=openai.OpenAIError("The api_key client option must be set either"), ) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_client_fails_because_openai_error_trace_no_exp( @@ -308,7 +308,7 @@ def call_openai(model: str = "gpt-3.5-turbo"): @patch("openai.resources.chat.Completions.create") -@patch("splunk_ao.logger.logger.LogStreams", side_effect=Exception("error")) +@patch("splunk_ao.logger.logger.AgentStreams", side_effect=Exception("error")) @patch("splunk_ao.logger.logger.Projects", side_effect=Exception("error")) @patch("splunk_ao.logger.logger.Traces") def test_galileo_api_client_transport_error_not_blocking_user_code( @@ -348,7 +348,7 @@ def call_openai(model: str = "gpt-3.5-turbo"): @patch("openai.resources.chat.Completions.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_openai_calls_in_active_trace( @@ -386,7 +386,7 @@ def test_openai_calls_in_active_trace( @patch("openai.resources.chat.Completions.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_chat_completions_multiple_messages( @@ -447,7 +447,7 @@ def test_chat_completions_multiple_messages( @patch("openai.resources.responses.Responses.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_basic_responses_api_call( @@ -488,7 +488,7 @@ def test_basic_responses_api_call( @patch("openai.resources.responses.Responses.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_responses_api_with_tools( @@ -551,7 +551,7 @@ def test_responses_api_with_tools( @patch("openai.resources.responses.Responses.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_responses_api_multiple_messages( @@ -613,7 +613,7 @@ def test_responses_api_multiple_messages( @patch("openai.resources.responses.Responses.create") -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") def test_responses_api_streaming( diff --git a/tests/test_openai_agents.py b/tests/test_openai_agents.py index fdc3e03a..8f25ef6d 100644 --- a/tests/test_openai_agents.py +++ b/tests/test_openai_agents.py @@ -71,7 +71,7 @@ async def homework_guardrail(ctx, agent, input_data): decode_compressed_response=True, record_mode=vcr.mode.NEW_EPISODES, ) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") async def test_complex_agent( @@ -102,7 +102,7 @@ async def test_complex_agent( decode_compressed_response=True, record_mode=vcr.mode.NEW_EPISODES, ) -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") async def test_simple_agent( @@ -196,7 +196,7 @@ def _find_tool_spans(spans): @mark.asyncio -@patch("splunk_ao.logger.logger.LogStreams") +@patch("splunk_ao.logger.logger.AgentStreams") @patch("splunk_ao.logger.logger.Projects") @patch("splunk_ao.logger.logger.Traces") async def test_pre_built_tools_multiple_types( diff --git a/tests/testutils/setup.py b/tests/testutils/setup.py index ba488a37..2866591e 100644 --- a/tests/testutils/setup.py +++ b/tests/testutils/setup.py @@ -8,7 +8,7 @@ from pydantic import BaseModel -from splunk_ao.log_streams import LogStream +from splunk_ao.agent_streams import AgentStream from splunk_ao.logger.logger import SplunkAOLogger from splunk_ao.projects import Project from splunk_ao.resources.models import ExperimentResponse, ProjectType @@ -214,7 +214,7 @@ def setup_mock_logstreams_client(mock_logstreams_client: Mock): now = datetime.datetime.now() mock_instance = mock_logstreams_client.return_value mock_instance.get = Mock( - return_value=LogStream( + return_value=AgentStream( LogStreamResponse( id="6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9b", project_id="6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9a", @@ -225,7 +225,7 @@ def setup_mock_logstreams_client(mock_logstreams_client: Mock): ) ) mock_instance.create = Mock( - return_value=LogStream( + return_value=AgentStream( LogStreamResponse( id="6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9b", project_id="6c4e3f7e-4a9a-4e7e-8c1f-3a9a3a9a3a9a",