Skip to content

Commit 9539122

Browse files
committed
refac
1 parent 19db873 commit 9539122

3 files changed

Lines changed: 96 additions & 12 deletions

File tree

backend/open_webui/models/chat_messages.py

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
from sqlalchemy import select, delete, func, cast, Integer, distinct
77
from sqlalchemy.ext.asyncio import AsyncSession
88
from open_webui.internal.db import Base, get_async_db_context
9-
from open_webui.utils.response import normalize_usage
9+
from open_webui.utils.response import merge_usage, normalize_usage
1010
from pydantic import BaseModel, ConfigDict
1111
from sqlalchemy import (
1212
JSON,
@@ -203,10 +203,8 @@ async def upsert_message(
203203
# Extract and normalize usage
204204
usage = get_usage(data)
205205
if usage:
206-
# Deep-merge: preserve existing keys not present in new data
207-
# This prevents background tasks (follow-ups, title, tags)
208-
# from accidentally clearing the primary response's token counts
209-
existing.usage = {**(existing.usage or {}), **usage}
206+
existing_usage = normalize_usage(existing.usage or {}) if existing.usage else {}
207+
existing.usage = existing_usage if usage == existing_usage else merge_usage(existing_usage, usage)
210208
existing.updated_at = now
211209
await db.commit()
212210
await db.refresh(existing)

backend/open_webui/utils/middleware.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -112,7 +112,7 @@
112112
)
113113
from open_webui.utils.payload import apply_system_prompt_to_body
114114
from open_webui.utils.plugin import load_function_module_by_id
115-
from open_webui.utils.response import normalize_usage
115+
from open_webui.utils.response import merge_usage, normalize_usage
116116
from open_webui.utils.sanitize import sanitize_code
117117
from open_webui.utils.task import (
118118
get_task_model_id,
@@ -4142,8 +4142,8 @@ async def flush_pending_delta_data(threshold: int = 0):
41424142

41434143
# Normalize and capture usage for DB persistence
41444144
if response_metadata.get('usage'):
4145-
response_metadata['usage'] = normalize_usage(response_metadata['usage'])
4146-
usage = response_metadata['usage']
4145+
usage = merge_usage(usage, response_metadata['usage'])
4146+
response_metadata['usage'] = usage
41474147

41484148
processed_data.update(response_metadata)
41494149
processed_data.pop('done', None)
@@ -4162,7 +4162,7 @@ async def flush_pending_delta_data(threshold: int = 0):
41624162
raw_usage = data.get('usage', {}) or {}
41634163
raw_usage.update(data.get('timings', {})) # llama.cpp
41644164
if raw_usage:
4165-
usage = normalize_usage(raw_usage)
4165+
usage = merge_usage(usage, raw_usage)
41664166
await event_emitter(
41674167
{
41684168
'type': 'chat:completion',
@@ -4260,9 +4260,9 @@ async def flush_pending_delta_data(threshold: int = 0):
42604260
current_response_tool_call['function']['name'] = delta_name
42614261

42624262
if delta_arguments:
4263-
current_response_tool_call['function']['arguments'] += (
4264-
delta_arguments
4265-
)
4263+
current_response_tool_call['function'][
4264+
'arguments'
4265+
] += delta_arguments
42664266

42674267
# Emit pending tool calls in real-time
42684268
if response_tool_calls:

backend/open_webui/utils/response.py

Lines changed: 86 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import json
2+
from numbers import Number
23
from uuid import uuid4
34

45
from open_webui.utils.misc import (
@@ -50,6 +51,91 @@ def normalize_usage(usage: dict) -> dict:
5051
return result
5152

5253

54+
USAGE_TOKEN_KEYS = {
55+
'input_tokens',
56+
'output_tokens',
57+
'total_tokens',
58+
'prompt_tokens',
59+
'completion_tokens',
60+
}
61+
62+
USAGE_COST_KEYS = {
63+
'cost',
64+
'total_cost',
65+
'input_cost',
66+
'output_cost',
67+
'prompt_cost',
68+
'completion_cost',
69+
}
70+
71+
USAGE_DETAIL_KEYS = {
72+
'prompt_tokens_details',
73+
'completion_tokens_details',
74+
'input_tokens_details',
75+
'output_tokens_details',
76+
}
77+
78+
79+
def _is_numeric_usage_value(value) -> bool:
80+
return isinstance(value, Number) and not isinstance(value, bool)
81+
82+
83+
def _merge_numeric_usage_map(current: dict | None, incoming: dict | None) -> dict:
84+
current = current or {}
85+
incoming = incoming or {}
86+
result = {**current, **incoming}
87+
88+
for key in set(current) | set(incoming):
89+
current_value = current.get(key, 0)
90+
incoming_value = incoming.get(key, 0)
91+
if isinstance(current_value, dict) or isinstance(incoming_value, dict):
92+
result[key] = _merge_numeric_usage_map(
93+
current_value if isinstance(current_value, dict) else {},
94+
incoming_value if isinstance(incoming_value, dict) else {},
95+
)
96+
elif _is_numeric_usage_value(current_value) or _is_numeric_usage_value(incoming_value):
97+
result[key] = (current_value if _is_numeric_usage_value(current_value) else 0) + (
98+
incoming_value if _is_numeric_usage_value(incoming_value) else 0
99+
)
100+
101+
return result
102+
103+
104+
def merge_usage(current: dict | None, incoming: dict | None) -> dict:
105+
"""
106+
Merge usage payloads from multiple model calls into one cumulative usage dict.
107+
108+
Token fields are additive; non-numeric metadata keeps the latest provider value.
109+
"""
110+
current_usage = normalize_usage(current or {}) if current else {}
111+
incoming_usage = normalize_usage(incoming or {}) if incoming else {}
112+
113+
if not incoming_usage:
114+
return current_usage
115+
if not current_usage:
116+
return incoming_usage
117+
118+
result = {**current_usage, **incoming_usage}
119+
120+
for key in USAGE_TOKEN_KEYS | USAGE_COST_KEYS:
121+
if key in current_usage or key in incoming_usage:
122+
current_value = current_usage.get(key, 0)
123+
incoming_value = incoming_usage.get(key, 0)
124+
if _is_numeric_usage_value(current_value) or _is_numeric_usage_value(incoming_value):
125+
result[key] = (current_value if _is_numeric_usage_value(current_value) else 0) + (
126+
incoming_value if _is_numeric_usage_value(incoming_value) else 0
127+
)
128+
129+
for key in USAGE_DETAIL_KEYS:
130+
if isinstance(current_usage.get(key), dict) or isinstance(incoming_usage.get(key), dict):
131+
result[key] = _merge_numeric_usage_map(
132+
current_usage.get(key) if isinstance(current_usage.get(key), dict) else {},
133+
incoming_usage.get(key) if isinstance(incoming_usage.get(key), dict) else {},
134+
)
135+
136+
return result
137+
138+
53139
def convert_ollama_tool_call_to_openai(tool_calls: list) -> list:
54140
openai_tool_calls = []
55141
for tool_call in tool_calls:

0 commit comments

Comments
 (0)