"""Adapter from the verified text Open Agent API to ProcessingTransport.""" from __future__ import annotations import hashlib import threading import time from typing import Any, Callable, Iterator, Mapping, NoReturn, Optional, Set from agent_integration.client import ( OpenAgentAPIError, OpenAgentProtocolError, OpenAgentTransportError, ) from agent_integration.service import AgentResponseError, OpenAgentService from arr_processing.agent_trace import ( AgentTraceSink, AgentTraceStoreError, capture_failure_record, project_trace_event, trace_run_id, ) from arr_processing.contracts import PROCESSING_REQUEST_VERSION, ProcessingRequest from arr_processing.errors import ProcessingTransportError from arr_processing.runner import RemoteRunSnapshot class OpenAgentProcessingTransport: """Uses only text Session/Run APIs; file resolution remains a runtime provider seam.""" def __init__( self, service: OpenAgentService, *, message_builder: Optional[Callable[[ProcessingRequest], str]] = None, trace_store: Optional[AgentTraceSink] = None, ) -> None: self._service = service self._message_builder = message_builder or (lambda request: request.message()) self._trace_store = trace_store self._trace_threads: Set[threading.Thread] = set() self._trace_threads_lock = threading.Lock() self._closed = False def submit(self, request: ProcessingRequest, idempotency_key: str) -> RemoteRunSnapshot: conversation_id = self.conversation_id(request.job_id) try: message = self._message_builder(request) metadata = { "contract_version": PROCESSING_REQUEST_VERSION, "processing_kind": "opera_daily", } if self._trace_store is not None: try: return self._submit_traced( request, conversation_id=conversation_id, message=message, idempotency_key=idempotency_key, metadata=metadata, ) except ( OpenAgentAPIError, OpenAgentTransportError, OpenAgentProtocolError, AgentResponseError, ) as error: self._raise_mapped(error, stream_submit=True) response = self._service.send_message( conversation_id, message, message_id=idempotency_key, external_subject_id=conversation_id, metadata=metadata, ) except ( OpenAgentAPIError, OpenAgentTransportError, OpenAgentProtocolError, AgentResponseError, ) as error: self._raise_mapped(error) return self._snapshot(response) def close(self, timeout_seconds: float = 2.0) -> None: """Best-effort drain before the shared HTTP client is closed at shutdown.""" if timeout_seconds < 0: raise ValueError("trace shutdown timeout must not be negative") with self._trace_threads_lock: self._closed = True threads = list(self._trace_threads) deadline = time.monotonic() + timeout_seconds for thread in threads: thread.join(max(0.0, deadline - time.monotonic())) def get(self, request: ProcessingRequest, remote_run_id: str) -> RemoteRunSnapshot: try: response = self._service.get_run( self.conversation_id(request.job_id), remote_run_id, ) except ( OpenAgentAPIError, OpenAgentTransportError, OpenAgentProtocolError, AgentResponseError, ) as error: self._raise_mapped(error) return self._snapshot(response, expected_run_id=remote_run_id) def cancel(self, request: ProcessingRequest, remote_run_id: str) -> RemoteRunSnapshot: try: response = self._service.cancel_run( self.conversation_id(request.job_id), remote_run_id, ) except ( OpenAgentAPIError, OpenAgentTransportError, OpenAgentProtocolError, AgentResponseError, ) as error: self._raise_mapped(error) return self._snapshot(response, expected_run_id=remote_run_id) @staticmethod def conversation_id(job_id: str) -> str: digest = hashlib.sha256(job_id.encode("utf-8")).hexdigest() return "arrproc_" + digest[:48] def _submit_traced( self, request: ProcessingRequest, *, conversation_id: str, message: str, idempotency_key: str, metadata: Mapping[str, Any], ) -> RemoteRunSnapshot: events = iter( self._service.stream_message( conversation_id, message, message_id=idempotency_key, external_subject_id=conversation_id, metadata=dict(metadata), include_trace=True, ) ) sequence = 0 for event in events: sequence += 1 self._capture_trace(request, event, sequence) remote_run_id = trace_run_id(event) if remote_run_id is None: continue self._start_trace_drain(request, events, sequence) return RemoteRunSnapshot(remote_run_id, "running") raise AgentResponseError("trace stream ended before run.started") def _start_trace_drain( self, request: ProcessingRequest, events: Iterator[Any], sequence: int, ) -> None: thread = threading.Thread( target=self._drain_trace, args=(request, events, sequence), name=f"arr-agent-trace-{request.attempt_no}", daemon=True, ) with self._trace_threads_lock: if self._closed: try: close = getattr(events, "close", None) if callable(close): close() finally: self._append_trace( request.job_id, capture_failure_record( attempt_no=request.attempt_no, sequence=sequence + 1, error=RuntimeError("transport closed"), ), ) return self._trace_threads.add(thread) try: thread.start() except RuntimeError as error: with self._trace_threads_lock: self._trace_threads.discard(thread) close = getattr(events, "close", None) if callable(close): close() self._append_trace( request.job_id, capture_failure_record( attempt_no=request.attempt_no, sequence=sequence + 1, error=error, ), ) def _drain_trace( self, request: ProcessingRequest, events: Iterator[Any], sequence: int, ) -> None: try: for event in events: sequence += 1 self._capture_trace(request, event, sequence) except Exception as error: # best-effort diagnostics must not affect the run self._append_trace( request.job_id, capture_failure_record( attempt_no=request.attempt_no, sequence=sequence + 1, error=error, ), ) finally: with self._trace_threads_lock: self._trace_threads.discard(threading.current_thread()) def _capture_trace(self, request: ProcessingRequest, event: Any, sequence: int) -> None: projected = project_trace_event( event, attempt_no=request.attempt_no, sequence=sequence, ) if projected is not None: self._append_trace(request.job_id, projected) def _append_trace(self, job_id: str, record: Mapping[str, Any]) -> None: if self._trace_store is None: return try: self._trace_store.append(job_id, record) except (AgentTraceStoreError, OSError, ValueError, TypeError): pass @staticmethod def _snapshot( response: Mapping[str, Any], expected_run_id: Optional[str] = None, ) -> RemoteRunSnapshot: if not isinstance(response, Mapping): raise ProcessingTransportError("PROCESSING_REMOTE_PROTOCOL_INVALID", retryable=False) run_id = response.get("run_id") status = response.get("status") final_content = response.get("final_content") if ( not isinstance(run_id, str) or not run_id or not isinstance(status, str) or not status or (expected_run_id is not None and run_id != expected_run_id) or (final_content is not None and not isinstance(final_content, str)) ): raise ProcessingTransportError("PROCESSING_REMOTE_PROTOCOL_INVALID", retryable=False) return RemoteRunSnapshot(run_id, status.lower(), final_content) @staticmethod def _raise_mapped( error: BaseException, *, stream_submit: bool = False, ) -> NoReturn: if isinstance(error, OpenAgentAPIError): if error.active_run_conflict: raise ProcessingTransportError( "PROCESSING_ACTIVE_RUN_CONFLICT", retryable=True, active_run_conflict=True, ) from None if stream_submit and (error.status_code == 408 or error.status_code >= 500): raise ProcessingTransportError( "PROCESSING_REMOTE_SUBMISSION_AMBIGUOUS", retryable=False, ) from None if error.retryable: code = ( "PROCESSING_REMOTE_RATE_LIMITED" if error.status_code == 429 else "PROCESSING_REMOTE_UNAVAILABLE" ) raise ProcessingTransportError(code, retryable=True) from None raise ProcessingTransportError("PROCESSING_REMOTE_REJECTED", retryable=False) from None if isinstance(error, OpenAgentTransportError): if stream_submit: raise ProcessingTransportError( "PROCESSING_REMOTE_SUBMISSION_AMBIGUOUS", retryable=False, ) from None raise ProcessingTransportError( "PROCESSING_REMOTE_UNAVAILABLE", retryable=True ) from None if stream_submit: raise ProcessingTransportError( "PROCESSING_REMOTE_SUBMISSION_AMBIGUOUS", retryable=False, ) from None raise ProcessingTransportError( "PROCESSING_REMOTE_PROTOCOL_INVALID", retryable=False ) from None