feat: rebuild as multi-user web agent
This commit is contained in:
@@ -0,0 +1,142 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from agent_platform.auth import UserIdentity
|
||||
from agent_platform.models import get_model_spec
|
||||
from agent_platform.runtime.loop import AgentLoop
|
||||
from agent_platform.runtime.tools import TOOL_METADATA
|
||||
from agent_platform.store import RuntimeStore
|
||||
|
||||
|
||||
def tool_call(call_id: str, name: str, arguments: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"id": call_id,
|
||||
"type": "function",
|
||||
"function": {"name": name, "arguments": json.dumps(arguments)},
|
||||
}
|
||||
|
||||
|
||||
class ScriptedProvider:
|
||||
def __init__(self, responses: list[dict[str, Any]]) -> None:
|
||||
self.responses = responses
|
||||
self.requests = []
|
||||
|
||||
async def complete(self, **kwargs):
|
||||
self.requests.append(kwargs)
|
||||
return self.responses.pop(0)
|
||||
|
||||
|
||||
class ParallelRegistry:
|
||||
def __init__(self) -> None:
|
||||
self.started = 0
|
||||
self.both_started = asyncio.Event()
|
||||
|
||||
def specs(self, *, read_only=False, allow_delegate=True):
|
||||
return [TOOL_METADATA["read_file"].openai_spec()]
|
||||
|
||||
async def execute(self, name, arguments, context):
|
||||
self.started += 1
|
||||
if self.started == 2:
|
||||
self.both_started.set()
|
||||
await asyncio.wait_for(self.both_started.wait(), timeout=1)
|
||||
return {"ok": True, "path": arguments["path"]}
|
||||
|
||||
|
||||
async def test_read_only_tool_calls_run_in_parallel(settings) -> None:
|
||||
first = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
tool_call("1", "read_file", {"path": "a"}),
|
||||
tool_call("2", "read_file", {"path": "b"}),
|
||||
],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
final = {"choices": [{"message": {"role": "assistant", "content": "done"}, "finish_reason": "stop"}]}
|
||||
store = RuntimeStore(settings.database_url)
|
||||
await store.initialize()
|
||||
registry = ParallelRegistry()
|
||||
loop = AgentLoop(ScriptedProvider([first, final]), registry, store, max_tool_output_chars=10_000)
|
||||
events = []
|
||||
|
||||
async def callback(event_type, payload):
|
||||
events.append(event_type)
|
||||
|
||||
try:
|
||||
answer = await loop.run(
|
||||
spec=get_model_spec("work-light"),
|
||||
messages=[{"role": "user", "content": "inspect both"}],
|
||||
identity=UserIdentity("u1", "", "", "user"),
|
||||
raw_user_jwt="jwt",
|
||||
chat_id="c1",
|
||||
callback=callback,
|
||||
)
|
||||
assert answer == "done"
|
||||
assert registry.started == 2
|
||||
assert events.count("tool.completed") == 2
|
||||
finally:
|
||||
await store.close()
|
||||
|
||||
|
||||
class SerialRegistry:
|
||||
def __init__(self) -> None:
|
||||
self.active = 0
|
||||
self.max_active = 0
|
||||
|
||||
def specs(self, *, read_only=False, allow_delegate=True):
|
||||
return [TOOL_METADATA["write_file"].openai_spec()]
|
||||
|
||||
async def execute(self, name, arguments, context):
|
||||
self.active += 1
|
||||
self.max_active = max(self.max_active, self.active)
|
||||
await asyncio.sleep(0.02)
|
||||
self.active -= 1
|
||||
return {"ok": True}
|
||||
|
||||
|
||||
async def test_mutating_tool_calls_are_serialized(settings) -> None:
|
||||
first = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
tool_call("1", "write_file", {"path": "a", "content": "a"}),
|
||||
tool_call("2", "write_file", {"path": "b", "content": "b"}),
|
||||
],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
final = {"choices": [{"message": {"role": "assistant", "content": "done"}, "finish_reason": "stop"}]}
|
||||
store = RuntimeStore(settings.database_url)
|
||||
await store.initialize()
|
||||
registry = SerialRegistry()
|
||||
loop = AgentLoop(ScriptedProvider([first, final]), registry, store, max_tool_output_chars=10_000)
|
||||
|
||||
async def callback(event_type, payload):
|
||||
return None
|
||||
|
||||
try:
|
||||
await loop.run(
|
||||
spec=get_model_spec("work-light"),
|
||||
messages=[{"role": "user", "content": "write both"}],
|
||||
identity=UserIdentity("u1", "", "", "user"),
|
||||
raw_user_jwt="jwt",
|
||||
chat_id="c1",
|
||||
callback=callback,
|
||||
)
|
||||
assert registry.max_active == 1
|
||||
finally:
|
||||
await store.close()
|
||||
Reference in New Issue
Block a user