feat: sync latest ARR implementation

This commit is contained in:
Wyndham ARR
2026-07-31 15:11:42 +08:00
parent d6f8a747fa
commit bf7939dd1a
185 changed files with 17527 additions and 2260 deletions

View File

@@ -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