Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 46 additions & 3 deletions lagent/adapters/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,9 @@
to the actual model backend, and quietly records the full conversation history (trajectories).

Key Features:
- **Direct Passthrough**: Calls OpenAI models using OpenAI format, and Anthropic models using Anthropic format. Returns matching formats identically without forced translation.
- **Direct Passthrough**: Calls OpenAI models using OpenAI format, and
Anthropic models using Anthropic format. Returns matching formats
identically without forced translation.
- **Trajectory Recording**: Merges and retains conversation turns into `_records`.
Interrupted or duplicate prefix traces are intelligently filtered when calling
`get_messages()`.
Expand Down Expand Up @@ -34,7 +36,7 @@
import re
import uuid
from collections import defaultdict
from typing import Any, Dict, List, Optional
from typing import Any, Dict, List, Optional, Union
from urllib.parse import urlparse

import aiohttp
Expand Down Expand Up @@ -94,6 +96,30 @@ def _is_lmdeploy_input_length_error(response_data: dict[str, Any]) -> bool:
_OPENCLAW_TOOL_CALL_ID_RE = re.compile(r'^call_?([0-9a-fA-F]{8,})$')


def _merge_anthropic_system_messages(request_data: dict[str, Any]) -> None:
messages = request_data.get('messages')
if not isinstance(messages, list):
return
system_contents = [message['content'] for message in messages if message['role'] == 'system']
if not system_contents:
return

if request_data.get('system') is not None:
system_contents.insert(0, request_data['system'])

system_blocks: list[dict[str, Any]] = []
for content in system_contents:
if system_blocks:
system_blocks.append({'type': 'text', 'text': '\n\n'})
if isinstance(content, str):
system_blocks.append({'type': 'text', 'text': content})
else:
system_blocks.extend(content)

request_data['system'] = system_blocks
request_data['messages'] = [message for message in messages if message['role'] != 'system']


def _extract_openai_reasoning_delta(delta: dict[str, Any]) -> Optional[str]:
for key in _OPENAI_REASONING_DELTA_KEYS:
value = delta.get(key)
Expand Down Expand Up @@ -292,6 +318,14 @@ class SessionClient:
real_api_key: The actual API key to use when forwarding requests.
real_base_url: The actual LLM API base URL to forward to.
port: Port to listen on. 0 means auto-assign.
client_timeout: Optional timeout for requests from this proxy to the
upstream model endpoint. Pass either ``aiohttp.ClientTimeout`` or
a dict accepted by that class, for example
``{"total": None, "sock_connect": 30, "sock_read": 3600}``.
``None`` preserves aiohttp's default behavior.
merge_anthropic_system_messages: Move ``system`` role messages from an
Anthropic request's ``messages`` list into its top-level ``system``
field before forwarding and recording the request.
"""

def __init__(
Expand All @@ -302,11 +336,17 @@ def __init__(
session_id: Optional[str] = None,
extra_body: Optional[dict] = None,
http_proxy: Optional[str] = None,
client_timeout: Optional[Union[aiohttp.ClientTimeout, dict]] = None,
merge_anthropic_system_messages: bool = False,
):
self.real_api_key = real_api_key
self.real_base_url = real_base_url.rstrip('/')
self.port = port
self.http_proxy = http_proxy
if isinstance(client_timeout, dict):
client_timeout = aiohttp.ClientTimeout(**client_timeout)
self.client_timeout = client_timeout
self.merge_anthropic_system_messages = merge_anthropic_system_messages
self.session_id = session_id or ctx_session_id.get() or os.getenv('XTUNER_SESSION_ID') or str(uuid.uuid4().int)
self.extra_body = extra_body or {}
self._records: Dict[str, List[Dict[str, list]]] = defaultdict(list)
Expand Down Expand Up @@ -392,6 +432,8 @@ async def _handle_request(self, request: web.Request) -> web.Response:
# "context_management: Extra inputs are not permitted").
for _f in _DROP_ANTHROPIC_BODY_FIELDS:
request_data.pop(_f, None)
if self.merge_anthropic_system_messages:
_merge_anthropic_system_messages(request_data)
request_body = json.dumps(request_data).encode('utf-8')

# By default we assume the incoming request is already in the target format
Expand Down Expand Up @@ -436,7 +478,8 @@ async def _handle_request(self, request: web.Request) -> web.Response:

is_stream = provider_request_data.get('stream', False) if provider_request_data else False

async with aiohttp.ClientSession() as client:
client_kwargs = {} if self.client_timeout is None else {'timeout': self.client_timeout}
async with aiohttp.ClientSession(**client_kwargs) as client:
async with client.request(
method=request.method,
url=target_url,
Expand Down
211 changes: 210 additions & 1 deletion tests/test_adapters/test_session_client.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import asyncio
import json
import logging
import os
Expand All @@ -9,6 +8,216 @@
from lagent.adapters.proxy import SessionClient


def _proxy(**kwargs):
return SessionClient(
real_api_key='EMPTY',
real_base_url='http://example.test/v1',
session_id='timeout-test',
**kwargs,
)


async def _capture_forwarded_request(monkeypatch, proxy, payload, path='v1/messages'):
captured = {}

class Request:
headers = {}
match_info = {'path': path}
method = 'POST'
query_string = ''

async def read(self):
return json.dumps(payload).encode('utf-8')

class ClientSession:
async def __aenter__(self):
return self

async def __aexit__(self, exc_type, exc, traceback):
return False

def request(self, **kwargs):
captured.update(kwargs)
raise RuntimeError('stop before sending the request')

monkeypatch.setattr(aiohttp, 'ClientSession', lambda **kwargs: ClientSession())
with pytest.raises(RuntimeError, match='stop before sending the request'):
await proxy._handle_request(Request())
return json.loads(captured['data'])


def test_client_timeout_accepts_aiohttp_object():
timeout = aiohttp.ClientTimeout(total=None, sock_connect=30, sock_read=3600)
proxy = _proxy(client_timeout=timeout)

assert proxy.client_timeout is timeout


def test_client_timeout_builds_from_serializable_dict():
proxy = _proxy(
client_timeout={
'total': None,
'sock_connect': 30,
'sock_read': 3600,
}
)

assert isinstance(proxy.client_timeout, aiohttp.ClientTimeout)
assert proxy.client_timeout.total is None
assert proxy.client_timeout.sock_connect == 30
assert proxy.client_timeout.sock_read == 3600


@pytest.mark.asyncio
@pytest.mark.parametrize('client_timeout', [None, {'sock_read': 3600}])
async def test_client_timeout_is_passed_to_client_session(monkeypatch, client_timeout):
proxy = _proxy(client_timeout=client_timeout)
client_kwargs = None

def create_client(**kwargs):
nonlocal client_kwargs
client_kwargs = kwargs
raise RuntimeError('stop before sending the request')

class Request:
headers = {}
match_info = {'path': 'v1/chat/completions'}
method = 'POST'
query_string = ''

async def read(self):
return b'{}'

monkeypatch.setattr(aiohttp, 'ClientSession', create_client)
with pytest.raises(RuntimeError, match='stop before sending the request'):
await proxy._handle_request(Request())

expected = {} if client_timeout is None else {'timeout': proxy.client_timeout}
assert client_kwargs == expected


@pytest.mark.asyncio
async def test_anthropic_system_messages_are_preserved_by_default(monkeypatch):
proxy = _proxy()
payload = {
'model': 'fake-model',
'max_tokens': 32,
'messages': [
{'role': 'user', 'content': 'question'},
{'role': 'system', 'content': 'project instruction'},
],
}

forwarded = await _capture_forwarded_request(monkeypatch, proxy, payload)

assert forwarded['messages'] == payload['messages']
assert 'system' not in forwarded


@pytest.mark.asyncio
async def test_anthropic_request_without_system_messages_is_unchanged(monkeypatch):
proxy = _proxy(merge_anthropic_system_messages=True)
payload = {
'model': 'fake-model',
'max_tokens': 32,
'system': 'base instruction',
'messages': [{'role': 'user', 'content': 'question'}],
}

forwarded = await _capture_forwarded_request(monkeypatch, proxy, payload)

assert forwarded['system'] == payload['system']
assert forwarded['messages'] == payload['messages']


@pytest.mark.asyncio
async def test_anthropic_system_messages_are_merged_before_forwarding_and_recording(monkeypatch):
proxy = _proxy(merge_anthropic_system_messages=True)
payload = {
'model': 'fake-model',
'max_tokens': 32,
'system': [
{
'type': 'text',
'text': 'base instruction',
'cache_control': {'type': 'ephemeral'},
}
],
'messages': [
{'role': 'user', 'content': 'question'},
{'role': 'system', 'content': 'project instruction'},
{'role': 'assistant', 'content': 'working'},
{
'role': 'system',
'content': [
{
'type': 'text',
'text': 'later instruction',
'cache_control': {'type': 'ephemeral'},
}
],
},
{'role': 'user', 'content': 'continue'},
],
}

forwarded = await _capture_forwarded_request(monkeypatch, proxy, payload)

assert forwarded['system'] == [
{
'type': 'text',
'text': 'base instruction',
'cache_control': {'type': 'ephemeral'},
},
{'type': 'text', 'text': '\n\n'},
{'type': 'text', 'text': 'project instruction'},
{'type': 'text', 'text': '\n\n'},
{
'type': 'text',
'text': 'later instruction',
'cache_control': {'type': 'ephemeral'},
},
]
assert forwarded['messages'] == [
{'role': 'user', 'content': 'question'},
{'role': 'assistant', 'content': 'working'},
{'role': 'user', 'content': 'continue'},
]

record = proxy._build_anthropic_record(
forwarded,
{'content': [{'type': 'text', 'text': 'answer'}]},
)
assert record is not None
messages, _ = record
assert messages[0] == {
'role': 'system',
'content': 'base instruction\n\nproject instruction\n\nlater instruction',
}
assert [message['role'] for message in messages] == ['system', 'user', 'assistant', 'user', 'assistant']


@pytest.mark.asyncio
async def test_anthropic_batch_request_without_top_level_messages_is_unchanged(monkeypatch):
proxy = _proxy(merge_anthropic_system_messages=True)
payload = {
'requests': [
{
'custom_id': 'request-1',
'params': {
'model': 'fake-model',
'max_tokens': 32,
'messages': [{'role': 'user', 'content': 'question'}],
},
}
]
}

forwarded = await _capture_forwarded_request(monkeypatch, proxy, payload, path='v1/messages/batches')

assert forwarded['requests'] == payload['requests']


def test_get_messages_normalizes_openclaw_tool_call_id_underscore_loss():
proxy = SessionClient(
real_api_key="EMPTY",
Expand Down