275 lines
9.9 KiB
Python
275 lines
9.9 KiB
Python
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()
|