from __future__ import annotations
import enum
import sys
from typing import Any, Callable, Mapping, Iterator
from agent_first_data.format import (
Event,
LogLevel,
OutputOptions,
json_error,
json_log,
json_progress,
json_result,
_format_json,
_format_yaml,
_format_plain,
validate_protocol_event,
)
class OutputFormat(enum.Enum):
JSON = "json"
YAML = "yaml"
PLAIN = "plain"
class OutputTo(enum.Enum):
SPLIT = "split"
STDOUT = "stdout"
STDERR = "stderr"
@classmethod
def parse(cls, value: str) -> OutputTo:
try:
return cls(value)
except ValueError:
raise ValueError(
f"unsupported --output-to {value!r}; expected split, stdout, or stderr"
)
class LogFilters:
def __init__(self, filters: list[str]) -> None:
self._filters = filters
def enabled(self, event: str) -> bool:
if not self._filters:
return False
if "all" in self._filters:
return True
event_lower = event.lower()
return any(event_lower.startswith(f) for f in self._filters)
def __bool__(self) -> bool:
return bool(self._filters)
def __iter__(self) -> Iterator[str]:
return iter(self._filters)
def __len__(self) -> int:
return len(self._filters)
def append(self, item: str) -> None:
self._filters.append(item)
def cli_parse_output(s: str) -> OutputFormat:
try:
return OutputFormat(s)
except ValueError:
raise ValueError(
f"invalid --output format {s!r}: expected json, yaml, or plain"
)
def cli_parse_log_filters(entries: list[str]) -> LogFilters:
out: list[str] = []
for entry in entries:
s = entry.strip().lower()
if s and s not in out:
out.append(s)
return LogFilters(out)
def render(value: Any, format: OutputFormat, *, options: OutputOptions | None = None) -> str:
if format is OutputFormat.YAML:
return _format_yaml(value, options=options)
if format is OutputFormat.PLAIN:
return _format_plain(value, options=options)
return _format_json(value, options=options)
class CliEmitter:
def __init__(
self,
writer: Any,
format: OutputFormat,
output_options: OutputOptions | None = None,
log_fields: Callable[[], Mapping[str, Any]] | None = None,
*,
diagnostic: Any | None = None,
) -> None:
self._writer = writer
self._diagnostic = diagnostic
self._format = format
self._output_options = output_options or OutputOptions()
self._terminal_emitted = False
self._log_fields_provider = log_fields
@classmethod
def stream(
cls,
writer: Any,
format: OutputFormat,
output_options: OutputOptions | None = None,
log_fields: Callable[[], Mapping[str, Any]] | None = None,
) -> CliEmitter:
return cls(writer, format, output_options, log_fields)
@classmethod
def finite_with(
cls,
result_writer: Any,
diagnostic: Any,
format: OutputFormat,
output_options: OutputOptions | None = None,
log_fields: Callable[[], Mapping[str, Any]] | None = None,
) -> CliEmitter:
return cls(result_writer, format, output_options, log_fields, diagnostic=diagnostic)
@classmethod
def finite(
cls,
format: OutputFormat,
output_options: OutputOptions | None = None,
log_fields: Callable[[], Mapping[str, Any]] | None = None,
) -> CliEmitter:
return cls.finite_with(sys.stdout, sys.stderr, format, output_options, log_fields)
@classmethod
def from_output_to(
cls,
selector: OutputTo,
format: OutputFormat,
output_options: OutputOptions | None = None,
log_fields: Callable[[], Mapping[str, Any]] | None = None,
) -> CliEmitter:
if selector is OutputTo.SPLIT:
return cls.finite_with(sys.stdout, sys.stderr, format, output_options, log_fields)
if selector is OutputTo.STDOUT:
return cls.stream(sys.stdout, format, output_options, log_fields)
if selector is OutputTo.STDERR:
return cls.stream(sys.stderr, format, output_options, log_fields)
raise ValueError(f"unsupported OutputTo selector: {selector!r}")
def with_log_fields(self, provider: Callable[[], Mapping[str, Any]]) -> CliEmitter:
self._log_fields_provider = provider
return self
def emit(self, event: Event | dict) -> None:
if isinstance(event, Event):
envelope = event.to_dict()
else:
envelope = event
validate_protocol_event(envelope, strict=False)
kind = envelope["kind"]
if kind in ("log", "progress"):
if self._terminal_emitted:
raise RuntimeError("cannot emit non-terminal event after terminal event")
elif kind in ("result", "error"):
if self._terminal_emitted:
raise RuntimeError("cannot emit duplicate terminal event")
else:
raise ValueError(f"unsupported event kind {kind!r}")
if kind == "log" and self._log_fields_provider is not None:
provider_fields = self._log_fields_provider()
log_payload = envelope.get("log")
if provider_fields and isinstance(log_payload, dict):
merged_log = dict(provider_fields)
merged_log.update(log_payload)
envelope["log"] = merged_log
use_diagnostic = kind != "result" and self._diagnostic is not None
sink = self._diagnostic if use_diagnostic else self._writer
sink.write(render(envelope, self._format, options=self._output_options) + "\n")
flush = getattr(sink, "flush", None)
if flush is not None:
flush()
if kind in ("result", "error"):
self._terminal_emitted = True
def emit_validated_value(self, value: Any) -> None:
try:
validate_protocol_event(value, strict=True)
except ValueError as e:
raise ValueError(f"emit_validated_value failed validation: {e}") from e
self.emit(value)
def emit_result(self, result: Any) -> None:
event = json_result(result).build()
self.emit(event)
def emit_error(self, code: str, message: str) -> None:
event = json_error(code, message).build()
self.emit(event)
def emit_progress(self, message: str) -> None:
if not message or not isinstance(message, str):
raise ValueError("message must be a non-empty string")
event = json_progress({"message": message}).build()
self.emit(event)
def emit_log(self, level: LogLevel | str, message: str) -> None:
if isinstance(level, str):
level = LogLevel(level)
if not message or not isinstance(message, str):
raise ValueError("message must be a non-empty string")
event = json_log({"level": level.value, "message": message}).build()
self.emit(event)
def finish(self, event: Event | dict, success_code: int) -> int:
try:
self.emit(event)
except BrokenPipeError:
return 0
except Exception:
return 4
return success_code
def finish_result(self, payload: Any) -> int:
return self.finish(json_result(payload).build(), 0)
def build_cli_version(version: str) -> dict:
return json_result({"code": "version", "version": version}).build().to_dict()
def cli_render_version(
name: str,
version: str,
format: OutputFormat | None = None,
) -> str:
rendered = (
f"{name} {version}" if format is None else render(build_cli_version(version), format)
)
return rendered.rstrip("\n") + "\n"
def cli_handle_version_or_continue(raw_args: list[str], name: str, version: str) -> str | None:
version_requested = False
output_format: OutputFormat | None = None
output_error: ValueError | None = None
i = 0
while i < len(raw_args):
arg = raw_args[i]
if arg == "--":
break
if not arg.startswith("-"):
break
if arg in ("--version", "-V"):
version_requested = True
i += 1
continue
if arg == "--json":
if output_format is not None and output_format is not OutputFormat.JSON:
output_error = ValueError(
"conflicting output formats: --json conflicts with previous output format"
)
else:
output_format = OutputFormat.JSON
i += 1
continue
if arg == "--output" or arg.startswith("--output="):
value: str | None
if arg.startswith("--output="):
value = arg.split("=", 1)[1]
step = 1
elif i + 1 < len(raw_args) and not raw_args[i + 1].startswith("-"):
value = raw_args[i + 1]
step = 2
else:
value = None
step = 1
if value is None:
output_error = ValueError(
"missing value for --output: expected json, yaml, or plain"
)
else:
try:
parsed_output = cli_parse_output(value)
if output_format is not None and output_format is not parsed_output:
output_error = ValueError(
f"conflicting output formats: --output {value} conflicts with previous output format"
)
else:
output_format = parsed_output
except ValueError as e:
output_error = e
i += step
continue
i += 1
if not version_requested:
return None
if output_error is not None:
raise output_error
return cli_render_version(name, version, output_format)
def build_cli_error(message: str, hint: str | None = None) -> dict | Event:
if not message:
message = "unspecified error"
event = json_error("cli_error", message)
if hint is not None:
event = event.hint(hint)
return event.build().to_dict()