feat: rebuild as multi-user web agent
This commit is contained in:
@@ -0,0 +1,274 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Annotated, Any
|
||||
|
||||
import httpx
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, Header, HTTPException, status
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
|
||||
from agent_platform.auth import UserIdentity, decode_openwebui_identity, verify_service_bearer
|
||||
from agent_platform.config import Settings, get_settings
|
||||
from agent_platform.models import get_model_spec, openai_model_list
|
||||
from agent_platform.runtime.loop import AgentLoop, tool_event_details
|
||||
from agent_platform.runtime.provider import ModelProvider
|
||||
from agent_platform.runtime.schemas import ChatCompletionRequest
|
||||
from agent_platform.runtime.tools import ToolRegistry
|
||||
from agent_platform.store import RuntimeStore
|
||||
|
||||
|
||||
def _sse_chunk(model: str, content: str = "", finish_reason: str | None = None) -> bytes:
|
||||
payload = {
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex}",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": int(time.time()),
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": ({"content": content} if content else {}),
|
||||
"finish_reason": finish_reason,
|
||||
}
|
||||
],
|
||||
}
|
||||
return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
|
||||
|
||||
|
||||
def _completion(model: str, content: str) -> dict[str, Any]:
|
||||
return {
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex}",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _public_sse_line(line: str, public_model: str) -> bytes:
|
||||
if not line.startswith("data:"):
|
||||
return f"{line}\n".encode()
|
||||
value = line.removeprefix("data:").strip()
|
||||
if not value or value == "[DONE]":
|
||||
return f"{line}\n".encode()
|
||||
try:
|
||||
payload = json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return f"{line}\n".encode()
|
||||
if isinstance(payload, dict) and "model" in payload:
|
||||
payload["model"] = public_model
|
||||
return f"data: {json.dumps(payload, ensure_ascii=False)}\n".encode()
|
||||
|
||||
|
||||
async def _public_model_stream(response: httpx.Response, public_model: str) -> AsyncIterator[bytes]:
|
||||
try:
|
||||
async for line in response.aiter_lines():
|
||||
yield _public_sse_line(line, public_model)
|
||||
finally:
|
||||
await response.aclose()
|
||||
|
||||
|
||||
def _extract_identity(
|
||||
settings: Settings,
|
||||
authorization: str | None,
|
||||
user_jwt: str | None,
|
||||
) -> UserIdentity:
|
||||
verify_service_bearer(authorization, settings.internal_provider_key)
|
||||
return decode_openwebui_identity(user_jwt, settings.openwebui_forward_jwt_secret)
|
||||
|
||||
|
||||
def create_app(
|
||||
settings: Settings | None = None,
|
||||
*,
|
||||
store: RuntimeStore | None = None,
|
||||
provider: ModelProvider | None = None,
|
||||
tools: ToolRegistry | None = None,
|
||||
) -> FastAPI:
|
||||
settings = settings or get_settings()
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
settings.validate_runtime()
|
||||
runtime_store = store or RuntimeStore(settings.database_url)
|
||||
await runtime_store.initialize()
|
||||
model_provider = provider or ModelProvider(settings)
|
||||
registry = tools or ToolRegistry(settings, runtime_store)
|
||||
app.state.store = runtime_store
|
||||
app.state.provider = model_provider
|
||||
app.state.tools = registry
|
||||
app.state.loop = AgentLoop(
|
||||
model_provider,
|
||||
registry,
|
||||
runtime_store,
|
||||
max_tool_output_chars=settings.max_tool_output_chars,
|
||||
)
|
||||
yield
|
||||
await registry.close()
|
||||
await model_provider.close()
|
||||
await runtime_store.close()
|
||||
|
||||
app = FastAPI(title="K1412 Agent Runtime", version="0.1.0", lifespan=lifespan)
|
||||
|
||||
@app.get("/health")
|
||||
async def health() -> dict:
|
||||
return {"status": "ok"}
|
||||
|
||||
@app.get("/v1/models")
|
||||
async def models(
|
||||
authorization: Annotated[str | None, Header(alias="Authorization")] = None,
|
||||
) -> dict:
|
||||
verify_service_bearer(authorization, settings.internal_provider_key)
|
||||
return openai_model_list()
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
async def chat_completions(
|
||||
body: ChatCompletionRequest,
|
||||
authorization: Annotated[str | None, Header(alias="Authorization")] = None,
|
||||
user_jwt: Annotated[str | None, Header(alias="X-OpenWebUI-User-Jwt")] = None,
|
||||
chat_id: Annotated[str | None, Header(alias="X-OpenWebUI-Chat-Id")] = None,
|
||||
message_id: Annotated[str | None, Header(alias="X-OpenWebUI-Message-Id")] = None,
|
||||
):
|
||||
identity = _extract_identity(settings, authorization, user_jwt)
|
||||
try:
|
||||
spec = get_model_spec(body.model)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
||||
|
||||
stable_chat_id = (chat_id or message_id or f"ephemeral-{uuid.uuid4().hex}").strip()
|
||||
selected_mode = await app.state.store.select_mode(identity.user_id, chat_id or "", spec.mode)
|
||||
if selected_mode == "work" and spec.mode == "chat":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail="This conversation has been upgraded to Work and cannot return to Chat.",
|
||||
)
|
||||
|
||||
if spec.mode == "chat":
|
||||
payload = body.model_dump(exclude_none=True)
|
||||
payload["model"] = spec.provider_model
|
||||
response = await app.state.provider.forward(payload)
|
||||
if body.stream:
|
||||
passthrough_headers = {
|
||||
key: value
|
||||
for key, value in response.headers.items()
|
||||
if key.lower() in {"cache-control", "x-request-id"}
|
||||
}
|
||||
return StreamingResponse(
|
||||
_public_model_stream(response, body.model),
|
||||
media_type=response.headers.get("content-type", "text/event-stream"),
|
||||
headers=passthrough_headers,
|
||||
)
|
||||
content = json.loads(await response.aread())
|
||||
if isinstance(content, dict) and "model" in content:
|
||||
content["model"] = body.model
|
||||
headers = {
|
||||
key: value
|
||||
for key, value in response.headers.items()
|
||||
if key.lower() in {"content-type", "cache-control", "x-request-id"}
|
||||
}
|
||||
await response.aclose()
|
||||
return JSONResponse(
|
||||
content=content,
|
||||
status_code=response.status_code,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
messages = [message.model_dump(exclude_none=True) for message in body.messages]
|
||||
if not body.stream:
|
||||
|
||||
async def ignore_event(_: str, __: dict[str, Any]) -> None:
|
||||
return None
|
||||
|
||||
answer = await app.state.loop.run(
|
||||
spec=spec,
|
||||
messages=messages,
|
||||
identity=identity,
|
||||
raw_user_jwt=user_jwt or "",
|
||||
chat_id=stable_chat_id,
|
||||
callback=ignore_event,
|
||||
)
|
||||
return _completion(body.model, answer)
|
||||
|
||||
async def work_stream() -> AsyncIterator[bytes]:
|
||||
queue: asyncio.Queue[tuple[str, Any]] = asyncio.Queue()
|
||||
|
||||
async def publish(event_type: str, payload: dict[str, Any]) -> None:
|
||||
details = tool_event_details(event_type, payload)
|
||||
if details:
|
||||
await queue.put(("content", details))
|
||||
|
||||
async def run_loop() -> None:
|
||||
try:
|
||||
answer = await app.state.loop.run(
|
||||
spec=spec,
|
||||
messages=messages,
|
||||
identity=identity,
|
||||
raw_user_jwt=user_jwt or "",
|
||||
chat_id=stable_chat_id,
|
||||
callback=publish,
|
||||
)
|
||||
await queue.put(("answer", answer))
|
||||
except Exception as exc:
|
||||
await queue.put(("error", str(exc)))
|
||||
finally:
|
||||
await queue.put(("done", None))
|
||||
|
||||
task = asyncio.create_task(run_loop())
|
||||
try:
|
||||
while True:
|
||||
kind, value = await queue.get()
|
||||
if kind == "done":
|
||||
break
|
||||
if kind == "error":
|
||||
yield _sse_chunk(body.model, f"\n\nWork 运行失败:{value}")
|
||||
continue
|
||||
if kind == "answer":
|
||||
text = str(value)
|
||||
for start in range(0, len(text), 240):
|
||||
yield _sse_chunk(body.model, text[start : start + 240])
|
||||
else:
|
||||
yield _sse_chunk(body.model, str(value))
|
||||
yield _sse_chunk(body.model, finish_reason="stop")
|
||||
yield b"data: [DONE]\n\n"
|
||||
finally:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
return StreamingResponse(
|
||||
work_stream(),
|
||||
media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
|
||||
)
|
||||
|
||||
@app.get("/v1/runs/{chat_id}")
|
||||
async def run_events(
|
||||
chat_id: str,
|
||||
authorization: Annotated[str | None, Header(alias="Authorization")] = None,
|
||||
user_jwt: Annotated[str | None, Header(alias="X-OpenWebUI-User-Jwt")] = None,
|
||||
) -> dict:
|
||||
identity = _extract_identity(settings, authorization, user_jwt)
|
||||
return {"items": await app.state.store.events_for_chat(identity.user_id, chat_id)}
|
||||
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
|
||||
|
||||
def run() -> None:
|
||||
uvicorn.run("agent_platform.runtime.app:app", host="0.0.0.0", port=8000) # noqa: S104
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run()
|
||||
Reference in New Issue
Block a user