feat: sync latest ARR implementation
This commit is contained in:
@@ -3,7 +3,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from typing import Any, Callable, Mapping, Optional
|
||||
import threading
|
||||
import time
|
||||
from typing import Any, Callable, Iterator, Mapping, NoReturn, Optional, Set
|
||||
|
||||
from agent_integration.client import (
|
||||
OpenAgentAPIError,
|
||||
@@ -11,6 +13,13 @@ from agent_integration.client import (
|
||||
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
|
||||
@@ -24,34 +33,79 @@ class OpenAgentProcessingTransport:
|
||||
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,
|
||||
self._message_builder(request),
|
||||
message,
|
||||
message_id=idempotency_key,
|
||||
external_subject_id=conversation_id,
|
||||
metadata={
|
||||
"contract_version": PROCESSING_REQUEST_VERSION,
|
||||
"processing_kind": "opera_daily",
|
||||
},
|
||||
metadata=metadata,
|
||||
)
|
||||
except (OpenAgentAPIError, OpenAgentTransportError, OpenAgentProtocolError, AgentResponseError) as error:
|
||||
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:
|
||||
except (
|
||||
OpenAgentAPIError,
|
||||
OpenAgentTransportError,
|
||||
OpenAgentProtocolError,
|
||||
AgentResponseError,
|
||||
) as error:
|
||||
self._raise_mapped(error)
|
||||
return self._snapshot(response, expected_run_id=remote_run_id)
|
||||
|
||||
@@ -61,7 +115,12 @@ class OpenAgentProcessingTransport:
|
||||
self.conversation_id(request.job_id),
|
||||
remote_run_id,
|
||||
)
|
||||
except (OpenAgentAPIError, OpenAgentTransportError, OpenAgentProtocolError, AgentResponseError) as error:
|
||||
except (
|
||||
OpenAgentAPIError,
|
||||
OpenAgentTransportError,
|
||||
OpenAgentProtocolError,
|
||||
AgentResponseError,
|
||||
) as error:
|
||||
self._raise_mapped(error)
|
||||
return self._snapshot(response, expected_run_id=remote_run_id)
|
||||
|
||||
@@ -70,6 +129,122 @@ class OpenAgentProcessingTransport:
|
||||
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],
|
||||
@@ -92,7 +267,11 @@ class OpenAgentProcessingTransport:
|
||||
return RemoteRunSnapshot(run_id, status.lower(), final_content)
|
||||
|
||||
@staticmethod
|
||||
def _raise_mapped(error: BaseException) -> None:
|
||||
def _raise_mapped(
|
||||
error: BaseException,
|
||||
*,
|
||||
stream_submit: bool = False,
|
||||
) -> NoReturn:
|
||||
if isinstance(error, OpenAgentAPIError):
|
||||
if error.active_run_conflict:
|
||||
raise ProcessingTransportError(
|
||||
@@ -100,6 +279,11 @@ class OpenAgentProcessingTransport:
|
||||
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"
|
||||
@@ -109,9 +293,19 @@ class OpenAgentProcessingTransport:
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user