# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Type stubs for ``nemo_relay.observability``."""

from __future__ import annotations

from dataclasses import dataclass, field
from typing import Literal

from nemo_relay import JsonObject, UnsupportedBehavior

MarkProjection = Literal["inherit", "event", "tool"]

@dataclass(slots=True)
class ConfigPolicy:
    unknown_component: UnsupportedBehavior = ...
    unknown_field: UnsupportedBehavior = ...
    unsupported_value: UnsupportedBehavior = ...
    def to_dict(self) -> JsonObject: ...

@dataclass(slots=True)
class AtofStreamSinkConfig:
    url: str = ...
    transport: Literal["http_post", "websocket", "ndjson"] = ...
    headers: dict[str, str] = field(default_factory=dict)
    header_env: dict[str, str] = field(default_factory=dict)
    timeout_millis: int = ...
    field_name_policy: Literal["preserve", "replace_dots"] = ...
    name: str | None = ...
    def to_dict(self) -> JsonObject: ...

@dataclass(slots=True)
class AtofConfig:
    enabled: bool = ...
    sinks: list[AtofFileSinkConfig | AtofStreamSinkConfig] | None = ...
    def to_dict(self) -> JsonObject: ...

@dataclass(slots=True)
class AtofFileSinkConfig:
    output_directory: str | None = ...
    filename: str | None = ...
    mode: Literal["append", "overwrite"] = ...
    def to_dict(self) -> JsonObject: ...

AtofEndpointConfig = AtofStreamSinkConfig

@dataclass(slots=True)
class S3StorageConfig:
    bucket: str = ...
    key_prefix: str | None = ...
    access_key_id: str | None = ...
    secret_access_key_var: str | None = ...
    session_token_var: str | None = ...
    region: str | None = ...
    endpoint_url: str | None = ...
    allow_http: bool | None = ...
    def to_dict(self) -> JsonObject: ...

@dataclass(slots=True)
class HttpStorageConfig:
    endpoint: str = ...
    headers: dict[str, str] = field(default_factory=dict)
    header_env: dict[str, str] = field(default_factory=dict)
    timeout_millis: int = ...
    def to_dict(self) -> JsonObject: ...

@dataclass(slots=True)
class AtifConfig:
    enabled: bool = ...
    agent_name: str = ...
    agent_version: str | None = ...
    model_name: str = ...
    tool_definitions: list[JsonObject] | None = ...
    extra: JsonObject | None = ...
    output_directory: str | None = ...
    filename_template: str = ...
    storage: list[S3StorageConfig | HttpStorageConfig] | None = ...
    def to_dict(self) -> JsonObject: ...

@dataclass(slots=True)
class OtlpConfig:
    enabled: bool = ...
    mark_projection: MarkProjection = ...
    mark_exclude_names: list[str] = ...
    transport: Literal["http_binary", "grpc"] = ...
    endpoint: str | None = ...
    headers: dict[str, str] = field(default_factory=dict)
    resource_attributes: dict[str, str] = field(default_factory=dict)
    service_name: str = ...
    service_namespace: str | None = ...
    service_version: str | None = ...
    instrumentation_scope: str | None = ...
    timeout_millis: int = ...
    attribute_mappings: list[dict[str, str]] = field(default_factory=list)
    def to_dict(self) -> JsonObject: ...

@dataclass(slots=True)
class ObservabilityConfig:
    version: int = ...
    atof: AtofConfig | None = ...
    atif: AtifConfig | None = ...
    opentelemetry: OtlpConfig | None = ...
    openinference: OtlpConfig | None = ...
    policy: ConfigPolicy = field(default_factory=ConfigPolicy)
    def to_dict(self) -> JsonObject: ...

OBSERVABILITY_PLUGIN_KIND: Literal["observability"]

@dataclass(slots=True)
class ComponentSpec:
    config: ObservabilityConfig | JsonObject
    enabled: bool = ...
    def to_dict(self) -> JsonObject: ...
