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__/__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__/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 38b0611c..2c3a92e9 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,4 @@ "splunk_ao_context", "start_session", ] + diff --git a/src/splunk_ao/log_stream.py b/src/splunk_ao/agent_stream.py similarity index 85% rename from src/splunk_ao/log_stream.py rename to src/splunk_ao/agent_stream.py index 81afefb4..a42af675 100644 --- a/src/splunk_ao/log_stream.py +++ b/src/splunk_ao/agent_stream.py @@ -9,7 +9,7 @@ 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.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, @@ -41,7 +41,10 @@ } -class LogStream(StateManagementMixin): +__all__ = ["AgentStream"] + + +class AgentStream(StateManagementMixin): """ Object-centric interface for Galileo log streams. @@ -63,12 +66,12 @@ class LogStream(StateManagementMixin): Examples -------- # Create a new log stream and persist it - log_stream = LogStream(name="Production Logs", project_name="My AI Project").create() + log_stream = AgentStream(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") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") - # LogStreams can also be created through Project instances + # AgentStreams can also be created through Project instances from splunk_ao.project import Project project = Project.get(name="My AI Project") @@ -76,7 +79,7 @@ class LogStream(StateManagementMixin): # Enable metrics on the log stream from splunk_ao.schema.metrics import SplunkAOMetrics - local_metrics = log_stream.enable_metrics([ + local_metrics = log_stream.enable_evaluators([ SplunkAOMetrics.correctness, SplunkAOMetrics.completeness, "context_relevance" @@ -97,15 +100,15 @@ class LogStream(StateManagementMixin): def __str__(self) -> str: """String representation of the log stream.""" - return f"LogStream(name='{self.name}', id='{self.id}', project_id='{self.project_id}')" + 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"LogStream(name='{self.name}', id='{self.id}', project_id='{self.project_id}', created_at='{self.created_at}')" + 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 LogStream instance locally. + 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. @@ -125,13 +128,13 @@ def __init__(self, name: str, *, project_id: str | None = None, project_name: st Examples -------- # Create by project ID - log_stream = LogStream(name="Production Logs", project_id="project-123") + log_stream = AgentStream(name="Production Logs", project_id="project-123") # Create by project name - log_stream = LogStream(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream(name="Production Logs", project_name="My AI Project") # Create using SPLUNK_AO_PROJECT environment variable - log_stream = LogStream(name="Production Logs") + log_stream = AgentStream(name="Production Logs") """ super().__init__() @@ -151,13 +154,13 @@ def __init__(self, name: str, *, project_id: str | None = None, project_name: st # Set initial state self._set_state(SyncState.LOCAL_ONLY) - def create(self) -> LogStream: + def create(self) -> AgentStream: """ Persist this log stream to the API. Returns ------- - LogStream: This log stream instance with updated attributes from the API. + AgentStream: This log stream instance with updated attributes from the API. Raises ------ @@ -166,7 +169,7 @@ def create(self) -> LogStream: Examples -------- - log_stream = LogStream(name="Production Logs", project_name="My AI Project").create() + log_stream = AgentStream(name="Production Logs", project_name="My AI Project").create() assert log_stream.is_synced() """ if not self.name: @@ -176,10 +179,10 @@ def create(self) -> LogStream: # 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") + 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 LogStream.get() and LogStream.list(). + # the contract of AgentStream.get() and AgentStream.list(). project_obj = _resolve_project(self.project_id, self.project_name) # Update project info from resolved project @@ -187,11 +190,10 @@ def create(self) -> LogStream: if self.project_name is None: self.project_name = project_obj.name - log_streams_service = LogStreams() - created_log_stream = log_streams_service.create( + agent_streams_svc = AgentStreams() + created_log_stream = agent_streams_svc.create( name=self.name, - project_id=self.project_id, - project_name=None, # Use resolved project_id + project_id=self.project_id, # always set by _resolve_project above ) # Update attributes from response @@ -206,31 +208,31 @@ def create(self) -> LogStream: # Set state to synced self._set_state(SyncState.SYNCED) - logger.info(f"LogStream.create: id='{self.id}' - completed") + 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"LogStream.create: name='{self.name}' - failed: {e}") + logger.error(f"AgentStream.create: name='{self.name}' - failed: {e}") raise @classmethod - def _create_empty(cls) -> LogStream: + def _create_empty(cls) -> AgentStream: """Internal constructor bypassing __init__ for API hydration.""" instance = cls.__new__(cls) - super(LogStream, instance).__init__() + super(AgentStream, instance).__init__() return instance @classmethod - def _from_api_response(cls, retrieved_log_stream: Any) -> LogStream: + def _from_api_response(cls, retrieved_log_stream: Any) -> AgentStream: """ - Factory method to create a LogStream instance from an API response. + Factory method to create a AgentStream 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. + AgentStream: A new AgentStream instance populated with the API data. """ instance = cls._create_empty() instance.created_at = retrieved_log_stream.created_at @@ -246,7 +248,7 @@ def _from_api_response(cls, retrieved_log_stream: Any) -> LogStream: return instance @classmethod - def get(cls, *, name: str, project_id: str | None = None, project_name: str | None = None) -> LogStream | None: + def get(cls, *, name: str, project_id: str | None = None, project_name: str | None = None) -> AgentStream | None: """ Get an existing log stream by name. @@ -259,7 +261,7 @@ def get(cls, *, name: str, project_id: str | None = None, project_name: str | No Returns ------- - Optional[LogStream]: The log stream if found, None otherwise. + Optional[AgentStream]: The log stream if found, None otherwise. Raises ------ @@ -268,24 +270,24 @@ def get(cls, *, name: str, project_id: str | None = None, project_name: str | No Examples -------- # Get by project name - log_stream = LogStream.get( + log_stream = AgentStream.get( name="Production Logs", project_name="My AI Project" ) # Get by project ID - log_stream = LogStream.get( + log_stream = AgentStream.get( name="Production Logs", project_id="project-123" ) # Get using SPLUNK_AO_PROJECT environment variable - log_stream = LogStream.get(name="Production Logs") + log_stream = AgentStream.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) + 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 @@ -302,7 +304,7 @@ def list( project_name: str | None = None, limit: Unset | int = 100, starting_token: Unset | int = 0, - ) -> list[LogStream]: + ) -> list[AgentStream]: """ List log streams for a project. @@ -319,7 +321,7 @@ def list( Returns ------- - List[LogStream]: A page of log streams for the project. + List[AgentStream]: A page of log streams for the project. Raises ------ @@ -328,24 +330,24 @@ def list( Examples -------- # List by project name - log_streams = LogStream.list(project_name="My AI Project") + log_streams = AgentStream.list(project_name="My AI Project") # List by project ID - log_streams = LogStream.list(project_id="project-123") + log_streams = AgentStream.list(project_id="project-123") # List using SPLUNK_AO_PROJECT environment variable - log_streams = LogStream.list() + log_streams = AgentStream.list() # Cap the number of returned log streams - log_streams = LogStream.list(project_name="My AI Project", limit=3) + log_streams = AgentStream.list(project_name="My AI Project", limit=3) # Fetch the next page - page_2 = LogStream.list(project_name="My AI Project", starting_token=100) + page_2 = AgentStream.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( + agent_streams_svc = AgentStreams() + retrieved_log_streams = agent_streams_svc.list( project_id=project_obj.id, limit=limit, starting_token=starting_token ) @@ -379,9 +381,9 @@ def refresh(self) -> 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) + 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") @@ -398,10 +400,10 @@ def refresh(self) -> None: # Set state to synced self._set_state(SyncState.SYNCED) - logger.debug(f"LogStream.refresh: id='{self.id}' - completed") + logger.debug(f"AgentStream.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}") + logger.error(f"AgentStream.refresh: id='{self.id}' - failed: {e}") raise def get_metrics(self) -> builtins.list[str]: @@ -418,11 +420,11 @@ def get_metrics(self) -> builtins.list[str]: Examples -------- - log_stream = LogStream.get(name="Production Logs", project_name="My Project") + 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"LogStream.get_metrics: id='{self.id}' - started") + 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( @@ -430,12 +432,12 @@ def get_metrics(self) -> builtins.list[str]: ) if settings is None or not hasattr(settings, "scorers"): - logger.info(f"LogStream.get_metrics: id='{self.id}' - no metrics enabled") + 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"LogStream.get_metrics: id='{self.id}' found {len(metric_names)} metrics - completed") + logger.info(f"AgentStream.get_metrics: id='{self.id}' found {len(metric_names)} metrics - completed") return metric_names def set_metrics( @@ -465,9 +467,9 @@ def set_metrics( Examples -------- - from splunk_ao import Metric, LogStream + from splunk_ao import Metric, AgentStream - log_stream = LogStream.get(name="Production Logs", project_name="My Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My Project") # Set metrics (replaces existing) log_stream.set_metrics([ @@ -477,19 +479,19 @@ def set_metrics( ]) """ 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) + 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_metrics(metrics) + result = log_stream.enable_evaluators(metrics) # Set state to synced after successful operation self._set_state(SyncState.SYNCED) - logger.info(f"LogStream.enable_metrics: id='{self.id}' - completed") + 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"LogStream.enable_metrics: id='{self.id}' - failed: {e}") + logger.error(f"AgentStream.enable_evaluators: id='{self.id}' - failed: {e}") raise def query( @@ -525,7 +527,7 @@ def query( -------- from splunk_ao.search import RecordType - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") # Query with column-based filters and sort results = log_stream.query( @@ -555,7 +557,7 @@ def query( 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") + 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 @@ -621,7 +623,7 @@ def get_spans( Examples -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") # Get spans with filters and sorting spans = log_stream.get_spans( @@ -641,7 +643,7 @@ def get_spans( if spans.has_next_page: more_spans = spans.next_page() """ - logger.debug(f"LogStream.get_spans: id='{self.id}' limit={limit} - started") + 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 ) @@ -674,7 +676,7 @@ def get_traces( Examples -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") # Get traces with filters traces = log_stream.get_traces( @@ -690,7 +692,7 @@ def get_traces( for trace in traces: print(trace["id"], trace["input"]) """ - logger.debug(f"LogStream.get_traces: id='{self.id}' limit={limit} - started") + 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 ) @@ -723,7 +725,7 @@ def get_sessions( Examples -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") # Get sessions with filters sessions = log_stream.get_sessions( @@ -739,7 +741,7 @@ def get_sessions( for session in sessions: print(session["id"], session["model"]) """ - logger.debug(f"LogStream.get_sessions: id='{self.id}' limit={limit} - started") + 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 ) @@ -779,7 +781,7 @@ def export_records( -------- from splunk_ao.search import RecordType - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") # Export records with filters for record in log_stream.export_records( @@ -801,7 +803,7 @@ def export_records( root_type = RECORD_TYPE_TO_ROOT_TYPE[record_type] logger.info( - f"LogStream.export_records: id='{self.id}' record_type='{record_type.value}' " + f"AgentStream.export_records: id='{self.id}' record_type='{record_type.value}' " f"export_format='{export_format.value}' - started" ) export_client = ExportClient() @@ -829,7 +831,7 @@ def context(self) -> Any: Examples -------- - log_stream = LogStream.get( + log_stream = AgentStream.get( name="Production Logs", project_name="My AI Project" ) @@ -876,7 +878,7 @@ def span_columns(self) -> ColumnCollection: Examples -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") columns = log_stream.span_columns # Access a specific column @@ -909,7 +911,7 @@ def session_columns(self) -> ColumnCollection: Examples -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") columns = log_stream.session_columns # Access a specific column @@ -943,7 +945,7 @@ def trace_columns(self) -> ColumnCollection: Examples -------- - log_stream = LogStream.get(name="Production Logs", project_name="My AI Project") + log_stream = AgentStream.get(name="Production Logs", project_name="My AI Project") columns = log_stream.trace_columns # Access a specific column @@ -963,7 +965,7 @@ def trace_columns(self) -> ColumnCollection: return ColumnCollection(columns) -# Import at end to avoid circular import (project.py imports LogStream) +# 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, diff --git a/src/splunk_ao/log_streams.py b/src/splunk_ao/agent_streams.py similarity index 84% rename from src/splunk_ao/log_streams.py rename to src/splunk_ao/agent_streams.py index 662a646c..2a6ebb6e 100644 --- a/src/splunk_ao/log_streams.py +++ b/src/splunk_ao/agent_streams.py @@ -20,7 +20,7 @@ logger = get_logger(__name__) -class LogStream(LogStreamResponse): +class AgentStream(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 @@ -48,21 +48,21 @@ class LogStream(LogStreamResponse): -------- ```python # Create a new log stream in a project - from splunk_ao.log_streams import create_log_stream + from splunk_ao.log_streams import create_agent_stream # Create by project ID - log_stream = create_log_stream(name="Production Logs", project_id="project-123") + log_stream = create_agent_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") + 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_log_stream - log_stream = get_log_stream(name="Production Logs", project_name="My AI Project") + 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_log_streams - log_streams = list_log_streams(project_name="My AI 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})") @@ -77,7 +77,7 @@ class LogStream(LogStreamResponse): ) # Enable metrics on a log stream - RECOMMENDED APPROACH - from splunk_ao.log_streams import enable_metrics + from splunk_ao.log_streams import enable_evaluators from splunk_ao.schema.metrics import SplunkAOMetrics # Set environment variables first @@ -85,15 +85,15 @@ class LogStream(LogStreamResponse): # export SPLUNK_AO_PROJECT="My AI Project" # Clean and simple - just pass the metrics! - local_metrics = enable_metrics([ + local_metrics = enable_evaluators([ SplunkAOMetrics.correctness, SplunkAOMetrics.completeness, "context_relevance" ]) # Alternative: Use explicit parameters - local_metrics = enable_metrics( - log_stream_name="Production Logs", + local_metrics = enable_evaluators( + agent_stream_name="Production Logs", project_name="My AI Project", metrics=["correctness", "completeness"] ) @@ -102,7 +102,7 @@ class LogStream(LogStreamResponse): def __init__(self, log_stream: None | LogStreamResponse = None): """ - Initialize a LogStream instance. + Initialize a AgentStream instance. Parameters ---------- @@ -122,18 +122,18 @@ def __init__(self, log_stream: None | LogStreamResponse = None): self.additional_properties = log_stream.additional_properties.copy() return - def enable_metrics( + 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 - LogStream object. The method leverages the log stream's existing project_id and id + 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 LogStream + 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." @@ -157,7 +157,7 @@ def enable_metrics( Raises ------ ValueError - - If this LogStream instance lacks required `id` or `project_id` attributes + - 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 @@ -168,15 +168,15 @@ def enable_metrics( Basic usage with built-in metrics: ```python - from splunk_ao.log_streams import LogStreams + from splunk_ao.log_streams import AgentStreams from splunk_ao.schema.metrics import SplunkAOMetrics # Get a log stream first - log_streams = LogStreams() + 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_metrics([ + local_metrics = log_stream.enable_evaluators([ SplunkAOMetrics.correctness, SplunkAOMetrics.completeness, "context_relevance", @@ -195,7 +195,7 @@ def enable_metrics( def custom_scorer(trace_or_span): return 0.75 # Your scoring logic - local_metrics = log_stream.enable_metrics([ + local_metrics = log_stream.enable_evaluators([ SplunkAOMetrics.correctness, "completeness", Metric(name="domain_relevance", version=3), @@ -211,12 +211,12 @@ def custom_scorer(trace_or_span): ----- **Requirements:** - - The LogStream instance must have valid `id` and `project_id` attributes - - These are automatically set when retrieving LogStream objects via LogStreams methods + - 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 LogStream object + - Use this method when you already have a AgentStream object - More intuitive than specifying project/log stream names again - Cleaner object-oriented design pattern """ @@ -227,7 +227,7 @@ def custom_scorer(trace_or_span): return local_metrics -class LogStreams: +class AgentStreams: config: SplunkAOConfig def __init__(self) -> None: @@ -236,12 +236,12 @@ def __init__(self) -> None: @overload def list( self, *, project_id: str, limit: Unset | int = 100, starting_token: Unset | int = 0 - ) -> builtins.list[LogStream]: ... + ) -> builtins.list[AgentStream]: ... @overload def list( self, *, project_name: str, limit: Unset | int = 100, starting_token: Unset | int = 0 - ) -> builtins.list[LogStream]: ... + ) -> builtins.list[AgentStream]: ... def list( self, @@ -250,7 +250,7 @@ def list( project_name: str | None = None, limit: Unset | int = 100, starting_token: Unset | int = 0, - ) -> builtins.list[LogStream]: + ) -> builtins.list[AgentStream]: """ Lists log streams. Exactly one of `project_id` or `project_name` must be provided. @@ -270,7 +270,7 @@ def list( Returns ------- - builtins.list[LogStream] + builtins.list[AgentStream] A page of log streams. Raises @@ -302,7 +302,7 @@ def list( 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] + 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. @@ -311,7 +311,7 @@ def list( # `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]: + 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, @@ -322,7 +322,7 @@ def _list_all(self, *, project_id: str) -> builtins.list[LogStream]: errors so mid-pagination failures aren't silently swallowed into a truncated result. """ - all_log_streams: builtins.list[LogStream] = [] + all_log_streams: builtins.list[AgentStream] = [] starting_token: int = 0 seen_tokens: set[int] = {starting_token} while True: @@ -337,7 +337,7 @@ def _list_all(self, *, project_id: str) -> builtins.list[LogStream]: 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) + 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: @@ -354,9 +354,9 @@ def _list_all(self, *, project_id: str) -> builtins.list[LogStream]: return all_log_streams @overload - def get(self, *, id: str, project_id: str | None = None, project_name: str | None = None) -> LogStream | None: ... + 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) -> LogStream | None: ... + def get(self, *, name: str, project_id: str | None = None, project_name: str | None = None) -> AgentStream | None: ... def get( self, *, @@ -364,7 +364,7 @@ def get( name: str | None = None, project_id: str | None = None, project_name: str | None = None, - ) -> LogStream | None: + ) -> AgentStream | None: """ Retrieves a log stream by id or name. @@ -381,7 +381,7 @@ def get( Returns ------- - Optional[LogStream] + Optional[AgentStream] The log stream if found, None otherwise. Raises @@ -412,7 +412,7 @@ def get( ) if not log_stream_response: return None - return LogStream(log_stream=log_stream_response) + return AgentStream(log_stream=log_stream_response) if name: for log_stream in self._list_all(project_id=project_id): @@ -421,11 +421,11 @@ def get( return None @overload - def create(self, name: str, *, project_id: str | None = None) -> LogStream: ... + def create(self, name: str, *, project_id: str | None = None) -> AgentStream: ... @overload - def create(self, name: str, *, project_name: str) -> LogStream: ... + def create(self, name: str, *, project_name: str) -> AgentStream: ... - def create(self, name: str, *, project_id: str | None = None, project_name: str | None = None) -> LogStream: + 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. @@ -440,7 +440,7 @@ def create(self, name: str, *, project_id: str | None = None, project_name: str Returns ------- - LogStream + AgentStream The created log stream. Raises @@ -474,12 +474,12 @@ def create(self, name: str, *, project_id: str | None = None, project_name: str if not response: raise ValueError("Unable to create log stream") - return LogStream(log_stream=response) + return AgentStream(log_stream=response) - def enable_metrics( + def enable_evaluators( self, *, - log_stream_name: str | None = None, + agent_stream_name: str | None = None, project_name: str | None = None, metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str], ) -> builtins.list[LocalMetricConfig]: @@ -489,12 +489,12 @@ def enable_metrics( 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 + The log stream name can be provided via the 'agent_stream_name' parameter or the SPLUNK_AO_LOG_STREAM environment variable. Parameters ---------- - log_stream_name : Optional[str], optional + 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. @@ -519,12 +519,12 @@ def enable_metrics( -------- ```python # Enable built-in metrics with explicit parameters - from splunk_ao.log_streams import LogStreams + from splunk_ao.log_streams import AgentStreams from splunk_ao.schema.metrics import SplunkAOMetrics - log_streams = LogStreams() - scorer_configs, local_metrics = log_streams.enable_metrics( - log_stream_name="Production Logs", + log_streams = AgentStreams() + scorer_configs, local_metrics = log_streams.enable_evaluators( + agent_stream_name="Production Logs", project_name="My AI Project", metrics=[ SplunkAOMetrics.correctness, @@ -536,7 +536,7 @@ def enable_metrics( # 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( + scorer_configs, local_metrics = log_streams.enable_evaluators( metrics=["correctness", "completeness"] ) @@ -547,8 +547,8 @@ 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 + 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), @@ -559,7 +559,7 @@ def custom_scorer(trace_or_span): """ # 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() + 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) @@ -567,9 +567,11 @@ def custom_scorer(trace_or_span): 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 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 '{log_stream_name}' not found in project '{project_obj.name}'") + raise ValueError(f"Log stream '{agent_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) @@ -581,9 +583,9 @@ def custom_scorer(trace_or_span): # -def get_log_stream( +def get_agent_stream( *, name: str | None = None, project_id: str | None = None, project_name: str | None = None -) -> LogStream | None: +) -> AgentStream | None: """ Retrieves a log stream by name. Exactly one of `project_id` or `project_name` must be provided. @@ -598,7 +600,7 @@ def get_log_stream( Returns ------- - Optional[LogStream] + Optional[AgentStream] The log stream if found, None otherwise. Raises @@ -610,16 +612,16 @@ def get_log_stream( httpx.TimeoutException If the request takes longer than Client.timeout. """ - return LogStreams().get(name=name, project_id=project_id, project_name=project_name) + return AgentStreams().get(name=name, project_id=project_id, project_name=project_name) # type: ignore[arg-type] -def list_log_streams( +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[LogStream]: +) -> builtins.list[AgentStream]: """ Lists log streams. Exactly one of `project_id` or `project_name` must be provided. @@ -639,7 +641,7 @@ def list_log_streams( Returns ------- - builtins.list[LogStream] + builtins.list[AgentStream] A page of log streams. Raises @@ -650,12 +652,12 @@ def list_log_streams( If the request takes longer than Client.timeout. """ - return LogStreams().list( + return AgentStreams().list( # type: ignore[call-overload] 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: +def create_agent_stream(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. @@ -666,7 +668,7 @@ def create_log_stream(name: str, project_id: str | None = None, project_name: st Returns ------- - LogStream + AgentStream The created project. Raises @@ -677,12 +679,12 @@ def create_log_stream(name: str, project_id: str | None = None, project_name: st If the request takes longer than Client.timeout. """ - return LogStreams().create(name=name, project_id=project_id, project_name=project_name) + return AgentStreams().create(name=name, project_id=project_id, project_name=project_name) # type: ignore[call-overload] -def enable_metrics( +def enable_evaluators( *, - log_stream_name: str | None = None, + agent_stream_name: str | None = None, project_name: str | None = None, metrics: builtins.list[SplunkAOMetrics | Metric | LocalMetricConfig | str], ) -> builtins.list[LocalMetricConfig]: @@ -702,11 +704,11 @@ def enable_metrics( 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) + The name of the log stream (used when agent_stream_name not provided) Parameters ---------- - log_stream_name : Optional[str], optional + 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 @@ -739,11 +741,11 @@ def enable_metrics( -------- ```python # Enable built-in metrics with explicit parameters - from splunk_ao.log_streams import enable_metrics + from splunk_ao.log_streams import enable_evaluators from splunk_ao.schema.metrics import SplunkAOMetrics - local_metrics = enable_metrics( - log_stream_name="Production Logs", + local_metrics = enable_evaluators( + agent_stream_name="Production Logs", project_name="My AI Project", metrics=[ SplunkAOMetrics.correctness, @@ -755,7 +757,7 @@ def enable_metrics( # 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"]) + local_metrics = enable_evaluators(metrics=["correctness", "completeness"]) # Enable custom and local metrics with environment variable fallbacks from splunk_ao.schema.metrics import Metric, LocalMetricConfig @@ -767,8 +769,8 @@ def response_length_scorer(trace_or_span): 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", + local_metrics = enable_evaluators( + agent_stream_name="Development Logs", metrics=[ SplunkAOMetrics.correctness, "toxicity", @@ -787,4 +789,4 @@ def response_length_scorer(trace_or_span): 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) + return AgentStreams().enable_evaluators(agent_stream_name=agent_stream_name, project_name=project_name, metrics=metrics) diff --git a/src/splunk_ao/metric.py b/src/splunk_ao/evaluator.py similarity index 83% rename from src/splunk_ao/metric.py rename to src/splunk_ao/evaluator.py index 3825e362..034ae5fb 100644 --- a/src/splunk_ao/metric.py +++ b/src/splunk_ao/evaluator.py @@ -19,7 +19,7 @@ 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.evaluators import Evaluators from splunk_ao.resources.api.data import ( create_code_scorer_version_scorers_scorer_id_version_code_post, create_scorers_post, @@ -40,7 +40,7 @@ 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.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 @@ -54,18 +54,18 @@ # - Configuration.code_validation_backoff_multiplier (env: SPLUNK_AO_CODE_VALIDATION_BACKOFF_MULTIPLIER) - default: 1.5 -class BuiltInMetrics: +class BuiltInEvaluators: """ Provides convenient access to built-in Galileo metrics (formerly "scorers"). Examples -------- - from splunk_ao.metric import Metric + from splunk_ao.metric import Evaluator # Access built-in metrics - Metric.metrics.correctness - Metric.metrics.completeness - Metric.metrics.toxicity + Evaluator.metrics.correctness + Evaluator.metrics.completeness + Evaluator.metrics.toxicity """ def __getattr__(self, name: str) -> SplunkAOMetrics: @@ -82,20 +82,20 @@ def __dir__(self) -> list[str]: # Backwards-compatible alias -BuiltInScorers = BuiltInMetrics +BuiltInScorers = BuiltInEvaluators -class Metric(StateManagementMixin, ABC): +class Evaluator(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) + - **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) Common Attributes ----------------- @@ -106,25 +106,25 @@ class Metric(StateManagementMixin, ABC): 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. + version (int | None): Evaluator version number. Class Attributes ---------------- - metrics (BuiltInMetrics): Access built-in Galileo metrics. + metrics (BuiltInEvaluators): Access built-in Galileo metrics. Examples -------- # 1. Use built-in Galileo scorers - from splunk_ao import Metric, SplunkAOMetric, LlmMetric, LocalMetric, LogStream + from splunk_ao import Evaluator, SplunkAOEvaluator, LlmEvaluator, LocalEvaluator, LogStream log_stream = LogStream.get(name="my-stream", project_name="my-project") log_stream.set_metrics([ - Metric.metrics.correctness, - Metric.metrics.completeness, + Evaluator.metrics.correctness, + Evaluator.metrics.completeness, ]) # 2. Create custom LLM metric - llm_metric = LlmMetric( + llm_metric = LlmEvaluator( name="response_quality", prompt="Rate the quality...", model="gpt-4o-mini", @@ -135,14 +135,14 @@ class Metric(StateManagementMixin, ABC): def my_scorer(trace_or_span): return 0.5 - local_metric = LocalMetric( + local_metric = LocalEvaluator( name="response_length", scorer_fn=my_scorer, ) """ # Class attribute for built-in metrics (preferred name) - metrics = BuiltInMetrics() + metrics = BuiltInEvaluators() # Backwards-compatible property for legacy name scorers = metrics @@ -167,7 +167,7 @@ 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. + Initialize a base Evaluator instance with common attributes. Args: name: The name of the metric. @@ -213,9 +213,9 @@ def _parse_output_type( return _map.get(output_type.lower(), default) @classmethod - def _create_metric_from_type(cls, scorer_type: ScorerTypes) -> Metric: + def _create_metric_from_type(cls, scorer_type: ScorerTypes) -> Evaluator: """ - Create the appropriate Metric subclass instance based on scorer_type. + 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. @@ -225,23 +225,23 @@ def _create_metric_from_type(cls, scorer_type: ScorerTypes) -> Metric: Returns ------- - Metric: An uninitialized instance of the appropriate subclass - (LlmMetric, CodeMetric, or SplunkAOMetric). + Evaluator: An uninitialized instance of the appropriate subclass + (LlmEvaluator, CodeEvaluator, or SplunkAOEvaluator). Examples -------- - instance = Metric._create_metric_from_type(ScorerTypes.LLM) - # Returns: LlmMetric instance + instance = Evaluator._create_metric_from_type(ScorerTypes.LLM) + # Returns: LlmEvaluator instance """ if scorer_type == ScorerTypes.LLM: - return LlmMetric.__new__(LlmMetric) + return LlmEvaluator.__new__(LlmEvaluator) if scorer_type == ScorerTypes.CODE: - return CodeMetric.__new__(CodeMetric) - # Default to SplunkAOMetric for built-in scorers (LUNA, PRESET, etc.) - return SplunkAOMetric.__new__(SplunkAOMetric) + 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) -> Metric | None: + def get(cls, *, id: str | None = None, name: str | None = None) -> Evaluator | None: """ Get an existing metric by ID or name. @@ -253,7 +253,7 @@ def get(cls, *, id: str | None = None, name: str | None = None) -> Metric | None Returns ------- - Optional[Metric]: The metric if found (SplunkAOMetric, LlmMetric, or CodeMetric), None otherwise. + Optional[Evaluator]: The metric if found (SplunkAOEvaluator, LlmEvaluator, or CodeEvaluator), None otherwise. Raises ------ @@ -262,10 +262,10 @@ def get(cls, *, id: str | None = None, name: str | None = None) -> Metric | None Examples -------- # Get by name - returns appropriate subclass - metric = Metric.get(name="factuality-checker") + metric = Evaluator.get(name="factuality-checker") # Get by ID - metric = Metric.get(id="abc-123-def") + 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") @@ -298,7 +298,7 @@ def get(cls, *, id: str | None = None, name: str | None = None) -> Metric | None @classmethod def list( cls, *, name_filter: str | None = None, scorer_types: list[ScorerTypes] | None = None - ) -> builtins.list[Metric]: + ) -> builtins.list[Evaluator]: """ List metrics with optional filtering. @@ -310,25 +310,25 @@ def list( Returns ------- - list[Metric]: List of metrics matching the criteria (with appropriate subclass types). + list[Evaluator]: List of metrics matching the criteria (with appropriate subclass types). Examples -------- # List all metrics - metrics = Metric.list() + metrics = Evaluator.list() # List LLM metrics only - metrics = Metric.list(scorer_types=[ScorerTypes.LLM]) + metrics = Evaluator.list(scorer_types=[ScorerTypes.LLM]) # List by name - metrics = Metric.list(name_filter="factuality") + metrics = Evaluator.list(name_filter="factuality") """ - logger.debug(f"Metric.list: name_filter='{name_filter}' types={scorer_types} - started") + 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"Metric.list: found {len(retrieved_scorers)} metrics - completed") + logger.debug(f"Evaluator.list: found {len(retrieved_scorers)} metrics - completed") - result: builtins.list[Metric] = [] + 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) @@ -344,7 +344,7 @@ 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()` + This is more efficient than calling `Evaluator.get(name=...).delete()` when you only need to delete and don't need the metric object. Args: @@ -357,19 +357,19 @@ def delete_by_name(cls, name: str) -> None: Examples -------- # Delete without retrieving first - Metric.delete_by_name("old-metric") + Evaluator.delete_by_name("old-metric") # Alternative (less efficient) - metric = Metric.get(name="old-metric") + metric = Evaluator.get(name="old-metric") metric.delete() """ - logger.info(f"Metric.delete_by_name: name='{name}' - started") + logger.info(f"Evaluator.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") + 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"Metric.delete_by_name: name='{name}' - failed: {e}") + logger.error(f"Evaluator.delete_by_name: name='{name}' - failed: {e}") raise def _populate_from_scorer_response(self, scorer_response: Any) -> None: @@ -410,8 +410,8 @@ def _populate_from_scorer_response(self, scorer_response: Any) -> None: cot_enabled=cot_enabled, ) - # LLM-specific attributes (only set if this is an LlmMetric) - if isinstance(self, LlmMetric): + # 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 @@ -432,8 +432,8 @@ def _populate_from_scorer_response(self, scorer_response: Any) -> None: 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): + # 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: @@ -447,7 +447,7 @@ def _populate_from_scorer_response(self, scorer_response: Any) -> None: self._sync_attrs(node_level=code_node_level, output_type=code_output_type) - def update(self, **kwargs: Any) -> Metric: + def update(self, **kwargs: Any) -> Evaluator: """ Update this metric's properties on the API. @@ -461,7 +461,7 @@ def update(self, **kwargs: Any) -> Metric: Returns ------- - Metric: This metric instance with updated attributes from the API. + Evaluator: This metric instance with updated attributes from the API. Raises ------ @@ -474,14 +474,14 @@ def update(self, **kwargs: Any) -> Metric: Examples -------- - metric = Metric.get(name="factuality-checker") + metric = Evaluator.get(name="factuality-checker") metric.update(name="new-name", description="Updated description") assert metric.is_synced() """ - if isinstance(self, LocalMetric): + 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("Metric ID is not set. Cannot update a local-only metric.") + 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: @@ -499,23 +499,23 @@ def update(self, **kwargs: Any) -> Metric: 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") + 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"Metric.update: id='{self.id}' - failed: {e}") + logger.error(f"Evaluator.update: id='{self.id}' - failed: {e}") raise if isinstance(response, HTTPValidationError): - raise APIError(f"Metric update validation error: {response.detail}") + raise APIError(f"Evaluator update validation error: {response.detail}") if response is None: - raise APIError(f"Metric update returned empty response for id '{self.id}'") + 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"Metric.update: id='{self.id}' - completed") + logger.info(f"Evaluator.update: id='{self.id}' - completed") return self def delete(self) -> None: @@ -531,24 +531,24 @@ def delete(self) -> None: Examples -------- - metric = Metric.get(name="factuality-checker") + metric = Evaluator.get(name="factuality-checker") metric.delete() """ - if isinstance(self, LocalMetric): + 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("Metric ID is not set. Cannot delete a local-only metric.") + raise ValueError("Evaluator 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) + 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"Metric.delete: id='{self.id}' - completed") + logger.info(f"Evaluator.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}") + logger.error(f"Evaluator.delete: id='{self.id}' - failed: {e}") raise def refresh(self) -> None: @@ -568,47 +568,47 @@ def refresh(self) -> None: metric.refresh() assert metric.is_synced() """ - if isinstance(self, LocalMetric): + 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("Metric ID is not set. Cannot refresh a local-only metric.") + raise ValueError("Evaluator ID is not set. Cannot refresh a local-only metric.") try: - logger.debug(f"Metric.refresh: id='{self.id}' - started") + 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"Metric with id '{self.id}' no longer exists") + 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"Metric.refresh: id='{self.id}' - completed") + logger.debug(f"Evaluator.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}") + logger.error(f"Evaluator.refresh: id='{self.id}' - failed: {e}") raise - def to_legacy_metric(self) -> LegacyMetric: + def to_legacy_metric(self) -> SchemaMetric: """ - Convert to legacy splunk_ao.schema.metrics.Metric format. + Convert to legacy splunk_ao.schema.metrics.Evaluator format. This enables backward compatibility with existing code that uses - the legacy Metric class. + the legacy Evaluator class. Returns ------- - LegacyMetric: Legacy metric object with name and version. + SchemaMetric: Legacy metric object with name and version. Examples -------- - metric = Metric.get(name="my-metric") + metric = Evaluator.get(name="my-metric") legacy = metric.to_legacy_metric() # Use with existing APIs """ - return LegacyMetric(name=self.name, version=self.version) + return SchemaMetric(name=self.name, version=self.version) def __str__(self) -> str: """String representation of the metric.""" @@ -623,11 +623,11 @@ def __repr__(self) -> str: # ============================================================================ -# Concrete Metric Types +# Concrete Evaluator Types # ============================================================================ -class LlmMetric(Metric): +class LlmEvaluator(Evaluator): """ LLM-based metric with custom prompt templates. @@ -652,7 +652,7 @@ class LlmMetric(Metric): Examples -------- # Create custom LLM metric with string model name - metric = LlmMetric( + metric = LlmEvaluator( name="response_quality", prompt=''' Rate the quality of this response on a scale of 1-10. @@ -674,7 +674,7 @@ class LlmMetric(Metric): # 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( + metric = LlmEvaluator( name="response_quality", prompt="Rate quality 1-10: {input} -> {output}", model=gpt_model, # Model object @@ -770,7 +770,7 @@ def __init__( # Handle output_type (accept string or enum) if isinstance(output_type, str): - self.output_type = Metric._parse_output_type(output_type, default=OutputTypeEnum.PERCENTAGE) + self.output_type = Evaluator._parse_output_type(output_type, default=OutputTypeEnum.PERCENTAGE) else: self.output_type = output_type or OutputTypeEnum.BOOLEAN @@ -778,13 +778,13 @@ def __init__( self.scorer_type = ScorerTypes.LLM - def create(self) -> LlmMetric: + def create(self) -> LlmEvaluator: """ Persist this LLM metric to the API. Returns ------- - LlmMetric: This metric instance with updated attributes from the API. + LlmEvaluator: This metric instance with updated attributes from the API. Raises ------ @@ -793,7 +793,7 @@ def create(self) -> LlmMetric: Examples -------- - metric = LlmMetric( + metric = LlmEvaluator( name="quality_check", prompt="Rate the quality...", model="gpt-4o-mini" @@ -801,10 +801,10 @@ def create(self) -> LlmMetric: assert metric.is_synced() """ try: - logger.info(f"LlmMetric.create: name='{self.name}' - started") + logger.info(f"LlmEvaluator.create: name='{self.name}' - started") - metrics_service = Metrics() - created_version = metrics_service.create_custom_llm_metric( + 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, @@ -829,21 +829,21 @@ def create(self) -> LlmMetric: # Refresh to get full scorer details self.refresh() - logger.info(f"LlmMetric.create: id='{self.id}' - completed") + 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"LlmMetric.create: name='{self.name}' - failed: {e}") + logger.error(f"LlmEvaluator.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})" + return f"LlmEvaluator(name='{self.name}', id='{self.id}', model='{self.model}', judges={self.judges})" -class CodeMetric(Metric): +class CodeEvaluator(Evaluator): r""" Code-based metric. @@ -859,11 +859,11 @@ class CodeMetric(Metric): Examples -------- # Get existing code metric - metric = Metric.get(name="my-code-metric") - assert isinstance(metric, CodeMetric) + metric = Evaluator.get(name="my-code-metric") + assert isinstance(metric, CodeEvaluator) # Create code metric with inline code - metric = CodeMetric( + metric = CodeEvaluator( name="custom_code_scorer", code="def scorer_fn(step_object):\\n return 1.0", description="Custom code-based scorer", @@ -873,7 +873,7 @@ class CodeMetric(Metric): ).create() # Load code from file - metric = CodeMetric( + metric = CodeEvaluator( name="custom_code_scorer", node_level=StepType.llm, ).load_code("./scorers/my_scorer.py").create() @@ -918,9 +918,9 @@ def __init__( self.required_metrics = required_metrics self.scorer_type = ScorerTypes.CODE - self.output_type = Metric._parse_output_type(output_type) + self.output_type = Evaluator._parse_output_type(output_type) - def load_code(self, code_file_path: str) -> CodeMetric: + def load_code(self, code_file_path: str) -> CodeEvaluator: """ Load code from a file into this metric instance. @@ -929,7 +929,7 @@ def load_code(self, code_file_path: str) -> CodeMetric: Returns ------- - CodeMetric: This metric instance with code loaded from the file (for chaining). + CodeEvaluator: This metric instance with code loaded from the file (for chaining). Raises ------ @@ -938,7 +938,7 @@ def load_code(self, code_file_path: str) -> CodeMetric: Examples -------- # Load code from file - metric = CodeMetric( + metric = CodeEvaluator( name="custom_code_scorer", node_level=StepType.llm, ).load_code("./scorers/my_scorer.py").create() @@ -985,11 +985,11 @@ def _validate_code(self, config: SplunkAOConfig) -> str: ) if validate_response is None: - logger.debug("CodeMetric._validate_code: No response from validate_code_scorer") + 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"CodeMetric._validate_code: task_id='{task_id}' - validation started") + 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() @@ -1005,11 +1005,11 @@ def _validate_code(self, config: SplunkAOConfig) -> str: ) if task_result is None: - logger.debug(f"CodeMetric._validate_code: No response for task_id='{task_id}'") + 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"CodeMetric._validate_code: task_id='{task_id}' - validation completed") + logger.debug(f"CodeEvaluator._validate_code: task_id='{task_id}' - validation completed") # Extract and validate the result result = task_result.result @@ -1043,7 +1043,7 @@ def _validate_code(self, config: SplunkAOConfig) -> str: ) logger.debug( - f"CodeMetric._validate_code: task_id='{task_id}' - pending " + 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) @@ -1051,7 +1051,7 @@ def _validate_code(self, config: SplunkAOConfig) -> str: else: raise ValueError(f"Unknown task status: {task_result.status}") - def create(self) -> CodeMetric: + def create(self) -> CodeEvaluator: r""" Persist this Code metric to the API. @@ -1061,7 +1061,7 @@ def create(self) -> CodeMetric: Returns ------- - CodeMetric: This metric instance with updated attributes from the API. + CodeEvaluator: This metric instance with updated attributes from the API. Raises ------ @@ -1071,7 +1071,7 @@ def create(self) -> CodeMetric: Examples -------- # Create with inline code - metric = CodeMetric( + metric = CodeEvaluator( name="custom_code_scorer", code="def scorer_fn(step_object):\\n return 1.0", node_level=StepType.llm, @@ -1079,7 +1079,7 @@ def create(self) -> CodeMetric: assert metric.is_synced() # Create by loading from file - metric = CodeMetric( + metric = CodeEvaluator( name="custom_code_scorer", node_level=StepType.llm, ).load_code("./scorers/my_scorer.py").create() @@ -1088,11 +1088,11 @@ def create(self) -> CodeMetric: # 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." + "Code is not set. Either pass 'code' to __init__() or use CodeEvaluator.load_code() to load from a file." ) try: - logger.info(f"CodeMetric.create: name='{self.name}' - started") + logger.info(f"CodeEvaluator.create: name='{self.name}' - started") config = SplunkAOConfig.get() @@ -1100,9 +1100,9 @@ def create(self) -> CodeMetric: 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") + logger.debug(f"CodeEvaluator.create: name='{self.name}' - validating code") validation_result = self._validate_code(config) - logger.debug(f"CodeMetric.create: name='{self.name}' - code validated successfully") + logger.debug(f"CodeEvaluator.create: name='{self.name}' - code validated successfully") # Step 2: Create the scorer scorer_request = CreateScorerRequest( @@ -1118,7 +1118,7 @@ def create(self) -> CodeMetric: 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") + logger.debug("CodeEvaluator.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 @@ -1135,7 +1135,7 @@ def create(self) -> CodeMetric: if created_version is None: logger.debug( - "CodeMetric.create: No response from create_code_scorer_version_scorers_scorer_id_version_code_post" + "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") @@ -1147,42 +1147,42 @@ def create(self) -> CodeMetric: # Refresh to get full scorer details self.refresh() - logger.info(f"CodeMetric.create: id='{self.id}' - completed") + 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"CodeMetric.create: name='{self.name}' - failed: {e}") + logger.error(f"CodeEvaluator.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}'')" + return f"CodeEvaluator(name='{self.name}', id='{self.id}'')" -class SplunkAOMetric(Metric): +class SplunkAOEvaluator(Evaluator): """ Built-in Galileo scorer metric. This metric type represents Galileo's built-in scorers like correctness, - completeness, toxicity, etc. Access these via `Metric.metrics`. + completeness, toxicity, etc. Access these via `Evaluator.metrics`. Examples -------- # Access built-in scorers - from splunk_ao import Metric, LogStream + from splunk_ao import Evaluator, 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, + Evaluator.metrics.correctness, + Evaluator.metrics.completeness, + Evaluator.metrics.toxicity, ]) # Or get by name - metric = Metric.get(name="correctness") - assert isinstance(metric, SplunkAOMetric) + metric = Evaluator.get(name="correctness") + assert isinstance(metric, SplunkAOEvaluator) """ def __init__( @@ -1201,7 +1201,7 @@ def __init__( # Galileo metrics can have various scorer types, set during population -class LocalMetric(Metric): +class LocalEvaluator(Evaluator): """ Local function-based metric. @@ -1224,7 +1224,7 @@ def response_length_scorer(trace_or_span): return min(len(trace_or_span.output) / 100.0, 1.0) return 0.0 - local_metric = LocalMetric( + local_metric = LocalEvaluator( name="response_length", scorer_fn=response_length_scorer, scorable_types=[StepType.llm], @@ -1300,7 +1300,7 @@ def to_local_metric_config(self) -> LocalMetricConfig: def my_scorer(trace): return 0.5 - metric = LocalMetric(name="test", scorer_fn=my_scorer) + metric = LocalEvaluator(name="test", scorer_fn=my_scorer) config = metric.to_local_metric_config() """ return LocalMetricConfig( @@ -1314,4 +1314,4 @@ 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})" + return f"LocalEvaluator(name='{self.name}', scorer_fn={fn_name})" diff --git a/src/splunk_ao/metrics.py b/src/splunk_ao/evaluators.py similarity index 96% rename from src/splunk_ao/metrics.py rename to src/splunk_ao/evaluators.py index 3aa864c7..cb4e08b5 100644 --- a/src/splunk_ao/metrics.py +++ b/src/splunk_ao/evaluators.py @@ -26,13 +26,13 @@ _logger = logging.getLogger(__name__) -class Metrics: +class Evaluators: config: SplunkAOConfig def __init__(self) -> None: self.config = SplunkAOConfig.get() - def delete_metric(self, name: str) -> None: + 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.") @@ -45,7 +45,7 @@ def delete_metric(self, name: str) -> None: if response is None: raise ValueError("Failed to delete metric.") - def create_custom_llm_metric( + def create_custom_llm_evaluator( self, name: str, user_prompt: str, @@ -148,7 +148,7 @@ def query( # Public functions -def create_custom_llm_metric( +def create_custom_llm_evaluator( name: str, user_prompt: str, node_level: StepType = StepType.llm, @@ -194,12 +194,12 @@ def create_custom_llm_metric( """ if tags is None: tags = [] - return Metrics().create_custom_llm_metric( + return Evaluators().create_custom_llm_evaluator( name, user_prompt, node_level, cot_enabled, model_name, num_judges, description, tags, output_type, ground_truth ) -def get_metrics( +def get_evaluators( project_id: str, start_time: datetime.datetime, end_time: datetime.datetime, @@ -235,7 +235,7 @@ def get_metrics( LogRecordsMetricsResponse A LogRecordsMetricsResponse object containing the query results, or None if the query fails. """ - return Metrics().query( + return Evaluators().query( project_id=project_id, start_time=start_time, end_time=end_time, @@ -247,7 +247,7 @@ def get_metrics( ) -def delete_metric(name: str) -> None: +def delete_evaluator(name: str) -> None: """ Deletes a metric by its name. @@ -256,4 +256,4 @@ def delete_metric(name: str) -> None: name The name of the metric to delete. """ - Metrics().delete_metric(name) + 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/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/project.py b/src/splunk_ao/project.py index 05ca34dd..d06d6874 100644 --- a/src/splunk_ao/project.py +++ b/src/splunk_ao/project.py @@ -16,9 +16,9 @@ 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 splunk_ao.prompt import Prompt logger = logging.getLogger(__name__) @@ -55,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) @@ -288,64 +288,57 @@ 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") """ 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 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 - - # Cap the number of returned log streams - log_streams = project.list_log_streams(limit=3) - - # Fetch the next page - page_2 = project.list_log_streams(starting_token=100) """ if self.id is None: - raise ValueError("Project ID is not set. Cannot list log streams for a local-only project.") + 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) - # Use the LogStream pattern to avoid duplication - return LogStream.list(project_id=self.id, limit=limit, starting_token=starting_token) def list_experiments(self) -> builtins.list[Experiment]: """ @@ -414,24 +407,25 @@ 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 experiments(self) -> builtins.list[Experiment]: @@ -890,8 +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 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",