Amazon Lex与LangChain集成:无法将session_context传入对话历史
问题描述
我已将对话历史保存到session_attributes['sessionContext']中,日志能看到该字段存储了完整对话历史,但LangChain Prompt中的{history}变量仅显示Lex的当前消息,需要将session_attributes中的对话历史正确传入Prompt。
问题代码
import all necessary pacakages logger = logging.getLogger() logger.setLevel(logging.DEBUG) def close(session_attributes, active_contexts, fulfillment_state, intent, message): response = { 'sessionState': { 'activeContexts':[{ 'name': 'intentContext', 'contextAttributes': active_contexts, 'timeToLive': { 'timeToLiveInSeconds': 600, 'turnsToLive': 1 } }], 'sessionAttributes': session_attributes, 'dialogAction': { 'type': 'Close', }, 'intent': intent, }, 'messages': [{'contentType': 'PlainText', 'content': message}] } return response def delegate(session_attributes, active_contexts, intent, message): return { 'sessionState': { 'activeContexts':[{ 'name': 'intentContext', 'contextAttributes': active_contexts, 'timeToLive': { 'timeToLiveInSeconds': 600, 'turnsToLive': 1 } }], 'sessionAttributes': session_attributes, 'dialogAction': { 'type': 'Delegate', }, 'intent': intent, }, 'messages': [{'contentType': 'PlainText', 'content': message}] } def initial_message(intent_name): response = { 'sessionState': { 'dialogAction': { 'type': 'ElicitSlot', 'slotToElicit': 'Location' if intent_name=='BookHotel' else 'PickUpCity' }, 'intent': { 'confirmationState': 'None', 'name': intent_name, 'state': 'InProgress' } } } return response # --- Helper Functions --- def try_ex(value): """ Call passed in function in try block. If KeyError is encountered return None. This function is intended to be used to safely access dictionary of the Slots section in the payloads. Note that this function would have negative impact on performance. """ if value is not None: return value['value']['interpretedValue'] else: return None def invoke_llm(query, session_history): endpoint_name = 'endpoint-name' region = 'us-east-1' kendra_index_id = 'kendra-index-id' print("invoke LLM session__history:: ", session_history ) class ContentHandler(ContentHandlerBase): content_type = "application/json" accepts = "application/json" def transform_input(self, prompt: str, model_kwargs: dict) -> bytes: input_str = json.dumps({"text_inputs": prompt, **model_kwargs}) return input_str.encode('utf-8') def transform_output(self, output: bytes) -> str: response_json = json.loads(output.read().decode("utf-8")) return response_json["generated_texts"][0] content_handler = ContentHandler() llm=SagemakerEndpoint( endpoint_name=endpoint_name, region_name=region, model_kwargs={"temperature":1e-10, "max_length": 500}, content_handler=content_handler ) retriever = KendraIndexRetriever(kendraindex=kendra_index_id, awsregion=region, return_source_documents=True) template = """ Use the following context (delimited by <ctx></ctx>) and the chat history (delimited by <hs></hs>) to answer the question: ------ <ctx> {context} </ctx> ------ <hs> {history} </hs> ------ {question} Answer: """ prompt = PromptTemplate( input_variables=["history", "context", "question"], template=template ) qa = RetrievalQA.from_chain_type( llm=llm, chain_type='stuff', retriever=retriever, verbose=True, chain_type_kwargs={ "verbose": True, "prompt": prompt, "memory": ConversationBufferMemory( memory_key="history", input_key="question", return_messages=True), } ) # chat_history = [] # while True: result = qa({'query':query, 'history': session_history}) # result = qa({'query':query, 'history': lex_conv_history}) response = result['result'] response = qa(query) print("Answer: ", response['result']) return response def lambda_handler(intent_request, context): print("input received: ", intent_request) logger.debug(intent_request) intent = intent_request['sessionState']['intent'] session_attributes = intent_request['sessionState']['sessionAttributes'] #print ("Session attributes -----",session_attributes) if 'sessionContext' not in session_attributes.keys(): print("First Execution") session_attributes['sessionContext'] = '' active_contexts = {} if intent['name']=='FallbackIntent': query = intent_request['inputTranscript']+session_attributes['sessionContext'] response = invoke_llm(query, session_attributes['sessionContext']) response_json = json.dumps({ 'Answer': response['result'], }) active_contexts['Query'] = response_json logger.debug('Answer from LLM={}'.format(response_json)) intent['confirmationState']="Confirmed" intent['state']="Fulfilled" session_attributes['sessionContext'] = session_attributes['sessionContext'] + ' ' + intent_request['inputTranscript'] + ' ' + response['result'] print ("History - - - ",session_attributes['sessionContext']) return close(session_attributes, active_contexts, 'Fulfilled', intent, response['result']) #confirmation_status = intent_request['sessionState']['intent']['confirmationState'] query = try_ex(intent_request['sessionState']['intent']['slots']['Query']) print("Question by user: ", query) if query or intent['name']=='FallbackIntent': response = invoke_llm(query, session_attributes['sessionContext']) # response = invoke_llm(query) response_json = json.dumps({ 'Answer': response['result'], }) logger.debug('Answer from LLM={}'.format(response_json)) intent['confirmationState']="Confirmed" intent['state']="Fulfilled" session_attributes['sessionContext'] = session_attributes['sessionContext'] + ' ' + intent_request['inputTranscript'] + ' ' + response['result'] print ("History - - - ",session_attributes['sessionContext']) return close(session_attributes, active_contexts, 'Fulfilled', intent, response['result'])
解决方案
问题出在两个核心点:
- 内置内存覆盖外部历史:
RetrievalQA配置了ConversationBufferMemory,它会自动管理对话历史,直接覆盖了你手动传入的session_history参数。 - 重复调用导致参数丢失:代码中先调用
qa({'query':query, 'history': session_history})得到正确结果,但随后又调用qa(query),这次调用未传入history,最终返回的是未携带历史的结果。
修改步骤:
- 移除内置内存配置:删除
chain_type_kwargs中的memory字段,让Prompt直接使用传入的history参数。 - 修正LLM调用逻辑:只保留一次正确调用,确保
history参数传递到Prompt中。 - 优化对话历史格式(可选):将
sessionContext的拼接格式调整为"用户:XXX\n助手:XXX",让LLM更易识别对话轮次。
修改后的invoke_llm函数:
def invoke_llm(query, session_history): endpoint_name = 'endpoint-name' region = 'us-east-1' kendra_index_id = 'kendra-index-id' print("invoke LLM session__history:: ", session_history ) class ContentHandler(ContentHandlerBase): content_type = "application/json" accepts = "application/json" def transform_input(self, prompt: str, model_kwargs: dict) -> bytes: input_str = json.dumps({"text_inputs": prompt, **model_kwargs}) return input_str.encode('utf-8') def transform_output(self, output: bytes) -> str: response_json = json.loads(output.read().decode("utf-8")) return response_json["generated_texts"][0] content_handler = ContentHandler() llm=SagemakerEndpoint( endpoint_name=endpoint_name, region_name=region, model_kwargs={"temperature":1e-10, "max_length": 500}, content_handler=content_handler ) retriever = KendraIndexRetriever(kendraindex=kendra_index_id, awsregion=region, return_source_documents=True) template = """ Use the following context (delimited by <ctx></ctx>) and the chat history (delimited by <hs></hs>) to answer the question: ------ <ctx> {context} </ctx> ------ <hs> {history} </hs> ------ {question} Answer: """ prompt = PromptTemplate( input_variables=["history", "context", "question"], template=template ) qa = RetrievalQA.from_chain_type( llm=llm, chain_type='stuff', retriever=retriever, verbose=True, chain_type_kwargs={ "verbose": True, "prompt": prompt, # 移除内置内存管理 } ) # 仅保留一次正确调用,传入完整历史参数 result = qa({'query': query, 'history': session_history}) response = result['result'] print("Answer: ", response) return result
在lambda_handler中更新对话历史时,建议调整格式:
# 替换原拼接逻辑,优化对话历史可读性 session_attributes['sessionContext'] += f"\n用户:{intent_request['inputTranscript']}\n助手:{response['result']}"
内容的提问来源于stack exchange,提问作者Utkarsh Rajput
相关产品推荐
相关产品推荐

