312 lines
11 KiB
Python
312 lines
11 KiB
Python
"""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
|