feat(chat): add configurable streaming output
This commit is contained in:
@@ -60,7 +60,7 @@ class ModelStreamService:
|
||||
model_id=model.id if model is not None else None,
|
||||
model_name=model_name,
|
||||
input_token=_rough_token_count(rag_result.prompt),
|
||||
chunks=_chunk_text(_mock_answer(rag_result)),
|
||||
chunks=_display_chunks(model, _mock_answer(rag_result)),
|
||||
)
|
||||
|
||||
if not rag_result.is_hit and not rag_result.allow_general_knowledge:
|
||||
@@ -68,7 +68,7 @@ class ModelStreamService:
|
||||
model_id=model.id if model is not None else None,
|
||||
model_name=model.model_name if model is not None else "no-hit",
|
||||
input_token=_rough_token_count(rag_result.prompt),
|
||||
chunks=_chunk_text(NO_HIT_ANSWER),
|
||||
chunks=_display_chunks(model, NO_HIT_ANSWER),
|
||||
)
|
||||
|
||||
if model is None:
|
||||
@@ -94,7 +94,7 @@ class ModelStreamService:
|
||||
model_id=model.id if model is not None else None,
|
||||
model_name=model_name,
|
||||
input_token=_rough_token_count(rag_result.prompt),
|
||||
chunks=_async_chunk_text(_mock_answer(rag_result)),
|
||||
chunks=_async_display_chunks(model, _mock_answer(rag_result)),
|
||||
)
|
||||
|
||||
if not rag_result.is_hit and not rag_result.allow_general_knowledge:
|
||||
@@ -102,7 +102,7 @@ class ModelStreamService:
|
||||
model_id=model.id if model is not None else None,
|
||||
model_name=model.model_name if model is not None else "no-hit",
|
||||
input_token=_rough_token_count(rag_result.prompt),
|
||||
chunks=_async_chunk_text(NO_HIT_ANSWER),
|
||||
chunks=_async_display_chunks(model, NO_HIT_ANSWER),
|
||||
)
|
||||
|
||||
if model is None:
|
||||
@@ -130,12 +130,12 @@ def _get_enabled_model(db: Session) -> ModelConfig | None:
|
||||
def _stream_configured_model(model: ModelConfig, rag_result: RagResult) -> Iterator[str]:
|
||||
api_type = model.api_type or "openai_compatible"
|
||||
if model.stream_enabled != 1:
|
||||
return _chunk_text(_call_configured_model(model, rag_result, allow_no_hit=rag_result.allow_general_knowledge))
|
||||
return iter((_call_configured_model(model, rag_result, allow_no_hit=rag_result.allow_general_knowledge),))
|
||||
if api_type == "anthropic_messages":
|
||||
return _stream_anthropic_messages(model, rag_result)
|
||||
if api_type == "openai_compatible":
|
||||
return _stream_openai_compatible_model(model, rag_result)
|
||||
return _chunk_text(_call_configured_model(model, rag_result, allow_no_hit=rag_result.allow_general_knowledge))
|
||||
return iter((_call_configured_model(model, rag_result, allow_no_hit=rag_result.allow_general_knowledge),))
|
||||
|
||||
|
||||
async def _stream_configured_model_async(model: ModelConfig, rag_result: RagResult) -> AsyncIterator[str]:
|
||||
@@ -144,8 +144,7 @@ async def _stream_configured_model_async(model: ModelConfig, rag_result: RagResu
|
||||
answer = await asyncio.to_thread(
|
||||
_call_configured_model, model, rag_result, allow_no_hit=rag_result.allow_general_knowledge
|
||||
)
|
||||
async for chunk in _async_chunk_text(answer):
|
||||
yield chunk
|
||||
yield answer
|
||||
return
|
||||
|
||||
if api_type == "anthropic_messages":
|
||||
@@ -161,8 +160,7 @@ async def _stream_configured_model_async(model: ModelConfig, rag_result: RagResu
|
||||
answer = await asyncio.to_thread(
|
||||
_call_configured_model, model, rag_result, allow_no_hit=rag_result.allow_general_knowledge
|
||||
)
|
||||
async for chunk in _async_chunk_text(answer):
|
||||
yield chunk
|
||||
yield answer
|
||||
|
||||
|
||||
def _stream_openai_compatible_model(model: ModelConfig, rag_result: RagResult) -> Iterator[str]:
|
||||
@@ -479,6 +477,19 @@ def _chunk_text(text: str, *, chunk_size: int = 12) -> Iterator[str]:
|
||||
async def _async_chunk_text(text: str, *, chunk_size: int = 12) -> AsyncIterator[str]:
|
||||
for index in range(0, len(text), chunk_size):
|
||||
yield text[index : index + chunk_size]
|
||||
await asyncio.sleep(0)
|
||||
|
||||
|
||||
def _display_chunks(model: ModelConfig | None, text: str) -> Iterator[str]:
|
||||
return _chunk_text(text) if model is None or model.stream_enabled == 1 else iter((text,))
|
||||
|
||||
|
||||
async def _async_display_chunks(model: ModelConfig | None, text: str) -> AsyncIterator[str]:
|
||||
if model is not None and model.stream_enabled != 1:
|
||||
yield text
|
||||
return
|
||||
async for chunk in _async_chunk_text(text):
|
||||
yield chunk
|
||||
|
||||
|
||||
def _string_or_empty(value: Any) -> str:
|
||||
|
||||
Reference in New Issue
Block a user