Skip to content

Commit a32d26e

Browse files
committed
refac
1 parent 989d5fd commit a32d26e

1 file changed

Lines changed: 21 additions & 28 deletions

File tree

backend/open_webui/models/chat_messages.py

Lines changed: 21 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,25 @@ def get_usage(data: dict) -> Optional[dict]:
4848
return normalize_usage(usage) if usage else None
4949

5050

51+
def _token_columns(dialect: str):
52+
"""Return (input_tokens, output_tokens) SQL column expressions.
53+
54+
Falls back to OpenAI-style keys (prompt_tokens / completion_tokens)
55+
when the normalized keys are absent.
56+
"""
57+
if dialect == 'sqlite':
58+
extract = lambda key: cast(func.json_extract(ChatMessage.usage, f'$.{key}'), Integer)
59+
elif dialect == 'postgresql':
60+
extract = lambda key: cast(func.json_extract_path_text(ChatMessage.usage, key), Integer)
61+
else:
62+
raise NotImplementedError(f'Unsupported dialect: {dialect}')
63+
64+
return (
65+
func.coalesce(extract('input_tokens'), extract('prompt_tokens')),
66+
func.coalesce(extract('output_tokens'), extract('completion_tokens')),
67+
)
68+
69+
5170
####################
5271
# ChatMessage DB Schema
5372
####################
@@ -343,20 +362,7 @@ async def get_token_usage_by_model(
343362
bind = await db.connection()
344363
dialect = bind.dialect.name
345364

346-
if dialect == 'sqlite':
347-
input_tokens = cast(func.json_extract(ChatMessage.usage, '$.input_tokens'), Integer)
348-
output_tokens = cast(func.json_extract(ChatMessage.usage, '$.output_tokens'), Integer)
349-
elif dialect == 'postgresql':
350-
input_tokens = cast(
351-
func.json_extract_path_text(ChatMessage.usage, 'input_tokens'),
352-
Integer,
353-
)
354-
output_tokens = cast(
355-
func.json_extract_path_text(ChatMessage.usage, 'output_tokens'),
356-
Integer,
357-
)
358-
else:
359-
raise NotImplementedError(f'Unsupported dialect: {dialect}')
365+
input_tokens, output_tokens = _token_columns(dialect)
360366

361367
stmt = select(
362368
ChatMessage.model_id,
@@ -404,20 +410,7 @@ async def get_token_usage_by_user(
404410
bind = await db.connection()
405411
dialect = bind.dialect.name
406412

407-
if dialect == 'sqlite':
408-
input_tokens = cast(func.json_extract(ChatMessage.usage, '$.input_tokens'), Integer)
409-
output_tokens = cast(func.json_extract(ChatMessage.usage, '$.output_tokens'), Integer)
410-
elif dialect == 'postgresql':
411-
input_tokens = cast(
412-
func.json_extract_path_text(ChatMessage.usage, 'input_tokens'),
413-
Integer,
414-
)
415-
output_tokens = cast(
416-
func.json_extract_path_text(ChatMessage.usage, 'output_tokens'),
417-
Integer,
418-
)
419-
else:
420-
raise NotImplementedError(f'Unsupported dialect: {dialect}')
413+
input_tokens, output_tokens = _token_columns(dialect)
421414

422415
stmt = select(
423416
ChatMessage.user_id,

0 commit comments

Comments
 (0)