Update
This commit is contained in:
@@ -33,7 +33,7 @@ from devplacepy.services.openai_gateway.vision import VisionAugmenter, VisionCac
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _fake_stream(data: dict, model: str):
|
||||
def _fake_stream(data: dict, model: str, include_usage: bool = False):
|
||||
chunk_id = data.get("id", f"chatcmpl-{uuid.uuid4().hex[:12]}")
|
||||
created = data.get("created", int(time.time()))
|
||||
out_model = data.get("model", model)
|
||||
@@ -72,6 +72,21 @@ def _fake_stream(data: dict, model: str):
|
||||
for i in range(0, len(content), 50):
|
||||
yield _chunk({"content": content[i : i + 50]})
|
||||
yield _chunk({}, finish="tool_calls" if tool_calls else "stop")
|
||||
if include_usage and data.get("usage"):
|
||||
yield (
|
||||
"data: "
|
||||
+ json.dumps(
|
||||
{
|
||||
"id": chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": out_model,
|
||||
"choices": [],
|
||||
"usage": data["usage"],
|
||||
}
|
||||
)
|
||||
+ "\n\n"
|
||||
)
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return gen()
|
||||
@@ -217,7 +232,7 @@ class GatewayRuntime:
|
||||
return resp, None, timing
|
||||
|
||||
async def handle_chat(
|
||||
self, body: dict, cfg: dict, owner: tuple, user_agent: str, log=None
|
||||
self, body: dict, cfg: dict, owner: tuple, user_agent: str, app_reference: str, log=None
|
||||
):
|
||||
log = log or (lambda message: None)
|
||||
overlay = chat_overlay(body.get("model"), cfg)
|
||||
@@ -242,6 +257,7 @@ class GatewayRuntime:
|
||||
owner=owner,
|
||||
pricing=pricing,
|
||||
context_map=context_map,
|
||||
app_reference=app_reference,
|
||||
)
|
||||
messages = await augmenter.augment_messages(client, messages)
|
||||
self.vision_calls += augmenter.calls
|
||||
@@ -252,14 +268,19 @@ class GatewayRuntime:
|
||||
requested = body.get("model")
|
||||
if cfg["gateway_force_model"] or not requested or requested == "molodetz":
|
||||
model = cfg["gateway_model"]
|
||||
else:
|
||||
elif overlay is not None:
|
||||
model = requested
|
||||
else:
|
||||
model = cfg["gateway_model"]
|
||||
log(f"requested model {requested!r} has no route, falling back to {model!r}")
|
||||
|
||||
stream = bool(body.get("stream"))
|
||||
include_usage = bool((body.get("stream_options") or {}).get("include_usage"))
|
||||
payload = dict(body)
|
||||
payload["model"] = model
|
||||
payload["messages"] = messages
|
||||
payload["stream"] = False
|
||||
payload.pop("stream_options", None)
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if cfg["gateway_api_key"]:
|
||||
@@ -287,6 +308,7 @@ class GatewayRuntime:
|
||||
"endpoint": "chat/completions",
|
||||
"model": model,
|
||||
"user_agent": user_agent,
|
||||
"app_reference": app_reference,
|
||||
**params,
|
||||
**timing,
|
||||
}
|
||||
@@ -372,14 +394,18 @@ class GatewayRuntime:
|
||||
log(f"chat POST -> 200 ({timing['upstream_latency_ms']:.0f}ms)")
|
||||
if stream:
|
||||
return StreamingResponse(
|
||||
_fake_stream(data, model),
|
||||
_fake_stream(data, model, include_usage),
|
||||
media_type="text/event-stream",
|
||||
headers=resp_headers,
|
||||
)
|
||||
return JSONResponse(content=data, headers=resp_headers)
|
||||
return Response(
|
||||
content=resp.content,
|
||||
media_type="application/json",
|
||||
headers=resp_headers,
|
||||
)
|
||||
|
||||
async def handle_embeddings(
|
||||
self, body: dict, cfg: dict, owner: tuple, user_agent: str, log=None
|
||||
self, body: dict, cfg: dict, owner: tuple, user_agent: str, app_reference: str, log=None
|
||||
):
|
||||
log = log or (lambda message: None)
|
||||
vision_cost = 0.0
|
||||
@@ -425,8 +451,11 @@ class GatewayRuntime:
|
||||
requested = body.get("model")
|
||||
if cfg["gateway_force_model"] or not requested or requested == "molodetz~embed":
|
||||
model = cfg["gateway_embed_model"]
|
||||
else:
|
||||
elif overlay is not None:
|
||||
model = requested
|
||||
else:
|
||||
model = cfg["gateway_embed_model"]
|
||||
log(f"requested embed model {requested!r} has no route, falling back to {model!r}")
|
||||
|
||||
payload = dict(body)
|
||||
payload["model"] = model
|
||||
@@ -457,6 +486,7 @@ class GatewayRuntime:
|
||||
"endpoint": "embeddings",
|
||||
"model": model,
|
||||
"user_agent": user_agent,
|
||||
"app_reference": app_reference,
|
||||
**params,
|
||||
**timing,
|
||||
}
|
||||
@@ -544,7 +574,7 @@ class GatewayRuntime:
|
||||
return JSONResponse(content=data, headers=resp_headers)
|
||||
|
||||
async def handle_images(
|
||||
self, body: dict, cfg: dict, owner: tuple, user_agent: str, log=None
|
||||
self, body: dict, cfg: dict, owner: tuple, user_agent: str, app_reference: str, log=None
|
||||
):
|
||||
log = log or (lambda message: None)
|
||||
overlay = image_overlay(body.get("model"), cfg)
|
||||
@@ -595,8 +625,11 @@ class GatewayRuntime:
|
||||
or requested in ("molodetz-img-small", "molodetz-img")
|
||||
):
|
||||
model = cfg["gateway_image_model"]
|
||||
else:
|
||||
elif overlay is not None:
|
||||
model = requested
|
||||
else:
|
||||
model = cfg["gateway_image_model"]
|
||||
log(f"requested image model {requested!r} has no route, falling back to {model!r}")
|
||||
|
||||
payload = dict(body)
|
||||
payload["model"] = model
|
||||
@@ -627,6 +660,7 @@ class GatewayRuntime:
|
||||
"endpoint": "images/generations",
|
||||
"model": model,
|
||||
"user_agent": user_agent,
|
||||
"app_reference": app_reference,
|
||||
**params,
|
||||
**timing,
|
||||
}
|
||||
@@ -717,6 +751,7 @@ class GatewayRuntime:
|
||||
cfg: dict,
|
||||
owner: tuple,
|
||||
user_agent: str,
|
||||
app_reference: str,
|
||||
log=None,
|
||||
):
|
||||
log = log or (lambda message: None)
|
||||
@@ -746,6 +781,7 @@ class GatewayRuntime:
|
||||
"endpoint": subpath,
|
||||
"model": cfg["gateway_model"],
|
||||
"user_agent": user_agent,
|
||||
"app_reference": app_reference,
|
||||
**timing,
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user