Support <think> tags (#117)

This commit is contained in:
Fabian Franz
2025-09-30 20:29:19 +08:00
committed by GitHub
parent 66cb51bb36
commit 8177876e5e
+17 -4
View File
@@ -268,6 +268,7 @@ class BedrockModel(BaseChatModel):
response = await self._invoke_bedrock(chat_request, stream=True) response = await self._invoke_bedrock(chat_request, stream=True)
message_id = self.generate_message_id() message_id = self.generate_message_id()
stream = response.get("stream") stream = response.get("stream")
self.think_emitted = False
async for chunk in self._async_iterate(stream): async for chunk in self._async_iterate(stream):
args = {"model_id": chat_request.model, "message_id": message_id, "chunk": chunk} args = {"model_id": chat_request.model, "message_id": message_id, "chunk": chunk}
stream_response = self._create_response_stream(**args) stream_response = self._create_response_stream(**args)
@@ -288,6 +289,7 @@ class BedrockModel(BaseChatModel):
# return an [DONE] message at the end. # return an [DONE] message at the end.
yield self.stream_response_to_bytes() yield self.stream_response_to_bytes()
self.think_emitted = False # Cleanup
except Exception as e: except Exception as e:
logger.error("Stream error for model %s: %s", chat_request.model, str(e)) logger.error("Stream error for model %s: %s", chat_request.model, str(e))
error_event = Error(error=ErrorMessage(message=str(e))) error_event = Error(error=ErrorMessage(message=str(e)))
@@ -634,6 +636,9 @@ class BedrockModel(BaseChatModel):
logger.warning( logger.warning(
"Unknown tag in message content " + ",".join(c.keys()) "Unknown tag in message content " + ",".join(c.keys())
) )
if message.reasoning_content:
message.content = f"<think>{message.reasoning_content}</think>{message.content}"
message.reasoning_content = None
response = ChatResponse( response = ChatResponse(
id=message_id, id=message_id,
@@ -702,11 +707,19 @@ class BedrockModel(BaseChatModel):
content=delta["text"], content=delta["text"],
) )
elif "reasoningContent" in delta: elif "reasoningContent" in delta:
# ignore "signature" in the delta.
if "text" in delta["reasoningContent"]: if "text" in delta["reasoningContent"]:
message = ChatResponseMessage( content = delta["reasoningContent"]["text"]
reasoning_content=delta["reasoningContent"]["text"], if not self.think_emitted:
) # Port of "content_block_start" with "thinking"
content = "<think>" + content
self.think_emitted = True
message = ChatResponseMessage(content=content)
elif "signature" in delta["reasoningContent"]:
# Port of "signature_delta"
if self.think_emitted:
message = ChatResponseMessage(content="\n </think> \n\n")
else:
return None # Ignore signature if no <think> started
else: else:
# tool use # tool use
index = chunk["contentBlockDelta"]["contentBlockIndex"] - 1 index = chunk["contentBlockDelta"]["contentBlockIndex"] - 1