Langchain对接Sagemaker流式响应配置失败,触发JSON解析错误
Langchain调用Sagemaker流式端点JSON解码错误排查与解决
问题场景
使用Langchain调用Sagemaker模型端点实现流式响应功能,运行代码后触发JSONDecodeError,报错指向langchain_community/llms/sagemaker_endpoint.py中的json.loads(line)语句。
报错分析
错误核心是流式响应的输出格式与Langchain的解析逻辑不匹配:
- Sagemaker流式端点的输出通常为**Server-Sent Events(SSE)**格式,每行以
data:开头,后跟随JSON内容,而非纯JSON字符串 - 当前
ContentHandler的transform_output方法直接读取整个输出并解析,未处理SSE格式前缀,导致解析空内容或无效字符串
解决方案
1. 适配SSE格式的输出解析
修改ContentHandler的transform_output方法,先去除每行的data: 前缀,再解析JSON;同时处理流式输出中的空行或结束标记(如data: [DONE])。
2. 匹配模型输入格式
根据Sagemaker端点部署的模型类型,调整transform_input中的请求结构,比如部分模型(如Llama 2)要求输入字段为prompt而非inputs。
3. 完善流式回调处理
在MyCustomHandler的on_llm_new_token中添加token打印,验证流式输出效果。
修改后的完整代码
from uuid import UUID from typing import Any, Dict, List from langchain.schema.output import LLMResult from langchain.callbacks.base import BaseCallbackHandler from langchain_community.llms.sagemaker_endpoint import LLMContentHandler from langchain_community.llms.sagemaker_endpoint import SagemakerEndpoint from langchain.prompts import PromptTemplate from langchain.chains.question_answering import load_qa_chain import json from langchain.docstore.document import Document from tenacity import RetryCallState # 示例文档与查询 example_doc_1 = """ Peter and Elizabeth took a taxi to attend the night party in the city. While in the party, Elizabeth collapsed and was rushed to the hospital. Since she was diagnosed with a brain injury, the doctor told Peter to stay besides her until she gets well. Therefore, Peter stayed with her at the hospital for 3 days without leaving. """ docs = [ Document(page_content=example_doc_1) ] query = "How long was Elizabeth hospitalized?" # 提示模板 prompt_template = """Use the following pieces of context to answer the question at the end. {context} Question: {question} Answer:""" PROMPT = PromptTemplate( template=prompt_template, input_variables=["context", "question"] ) # 自定义回调处理器 class MyCustomHandler(BaseCallbackHandler): def on_llm_new_token(self, token: str, **kwargs) -> None: # 实时打印流式输出的token print(token, end="", flush=True) def on_llm_error(self, error: BaseException, *, run_id: UUID, parent_run_id: UUID | None = None, **kwargs: Any) -> Any: print("\nLLM Error:", error) return super().on_llm_error(error, run_id=run_id, parent_run_id=parent_run_id, **kwargs) # 内容处理器(适配Sagemaker流式输出) class ContentHandler(LLMContentHandler): content_type = "application/json" accepts = "application/json" def transform_input(self, prompt: str, model_kwargs: Dict) -> bytes: # 根据模型调整输入结构,此处以Llama 2为例 input_dict = { "prompt": prompt, **model_kwargs, "stream": True } return json.dumps(input_dict).encode('utf-8') def transform_output(self, output: bytes) -> str: # 处理SSE格式输出:去除data前缀,跳过结束标记 output_str = output.read().decode("utf-8").strip() if not output_str or output_str == "[DONE]": return "" # 移除SSE格式的data:前缀 if output_str.startswith("data: "): output_str = output_str[6:] try: response_json = json.loads(output_str) # 根据模型输出结构调整,此处适配Llama 2的generation字段 return response_json.get("generation", "") except json.JSONDecodeError: return "" content_handler = ContentHandler() callback_f = MyCustomHandler() # 初始化Sagemaker端点LLM llm = SagemakerEndpoint( streaming=True, endpoint_name="your-endpoint-name", # 替换为实际端点名称 region_name="your-region", # 替换为实际区域 model_kwargs={"temperature": 0.7, "max_new_tokens": 400}, content_handler=content_handler, callbacks=[callback_f], verbose=True ) # 加载QA链并调用 chain = load_qa_chain( llm=llm, prompt=PROMPT, ) chain.invoke({"question": query, "input_documents": docs})
关键修改说明
- ContentHandler.transform_output:新增SSE格式处理逻辑,去除
data:前缀,跳过结束标记[DONE],避免解析无效内容 - ContentHandler.transform_input:将输入字段从
inputs改为prompt(适配Llama类模型,若你的模型使用inputs则改回) - MyCustomHandler.on_llm_new_token:添加token实时打印,直观验证流式输出效果
内容的提问来源于stack exchange,提问作者Mohamed Anser Ali
相关产品推荐
相关产品推荐

