您的RAG智能体在接收一条消息后就会忘记所有内容——如何用Databricks解决此问题……
您的RAG智能体在接收一条消息后就会忘记所有内容——如何用Databricks解决此问题……包括合同、校验功能以及为团队准备的即插即用代码模块。
本指南将逐步构建从原材料到可运行系统的完整流程,适用于“您的RAG智能体在接收一条消息后就忘记所有内容——我如何用Databricks Lakebase解决此问题”这一场景。重点在于可操作的步骤、明确的检查点,以及可直接放入代码仓库的代码,无需猜测其用途。 在修改代码之前,应先明确输入参数、各步骤的负责人以及完成标准。操作人员应能够从已知的检查点重新运行相应步骤,而无需推测隐藏的状态。 可将此阶段视为输入与经过验证的输出之间的契约。为相关成果命名,定义成功标准,杜绝无声的半完成状态。
架构设计
在研究《架构》时,首先写下相关契约:所需的输入参数、成功信号以及部分失败时的处理方式。这样的清单能确保后续的代码修改保持一致性。 在功能结果旁记录执行时间以及令牌或查询成本。提前了解成本情况,可避免在代码从演示环境转向共享环境时出现意外费用。 在调整提示词之前,先使用固定的问题集测试召回率。仅仅更换提示词往往无法改善较差的检索效果。
PDFs (text + images + diagrams)
↓ ai_parse_document() (Version 2.0)
Parsed elements (text, tables, figure descriptions)
↓ RecursiveCharacterTextSplitter
Chunks table (Delta, with Change Data Feed)
↓ Delta Sync + GTE-Large
Vector Search Index
↓ VectorSearchRetrieverTool
LangChain Agent + PostgresSaver
↓ ↓
LLM (Foundation Model) Lakebase (conversation memory)
↓
Model Serving Endpoint (MLflow ResponsesAgent)
步骤1:使用ai_parse_document()解析文档
在执行第一步:使用 ai_parse_document() 解析文档时,首先列出相关规范:所需输入、成功信号以及部分失败时的处理方式。这样的清单能确保后续的代码修改保持一致性。
将配置信息置于应用程序代码之外。环境文件、密钥存储以及功能开关应集中存放,以便操作人员无需查看整个系统结构即可进行审计。
在调整提示词之前,先使用固定的问题集来衡量召回率。仅仅更换提示词往往无法解决检索效果不佳的问题。
from pyspark.sql.functions import expr
# Volume path where PDFs are stored
docs_path = "/Volumes/<YOUR_CATALOG>/<YOUR_SCHEMA>/source_docs/"
# Read all files as binary
docs_df = spark.read.format("binaryFile").load(docs_path)
# Parse each document using ai_parse_document v2.0
parsed_df = docs_df.withColumn(
"parsed_content",
expr(f"""ai_parse_document(content, map(
"version", "2.0",
"imageOutputPath", "{docs_path}/parsed_images/"
))""")
)
# Drop binary content (too large to display)
parsed_df = parsed_df.drop("content")
# Save to Delta table
output_table = "<YOUR_CATALOG>.<YOUR_SCHEMA>.docs_parsed"
parsed_df.write.format("delta").mode("overwrite").saveAsTable(output_table)
print(f"✅ Parsed results saved to: {output_table}")
第二步:清洗、转换与分块
在执行第2步:清洗、转换和分块时,首先需明确相关约定:所需输入、成功信号以及部分失败时的处理方式。这样的清单能确保后续的代码修改保持一致性。 同时记录正常流程和异常恢复流程。重试机制、人工审核环节以及死信处理都是产品功能的一部分,而非后续的优化工作。 在调整提示词之前,需先用固定的问题集来衡量检索效果。仅仅更换提示词很难解决检索能力薄弱的问题。 在执行第2步:清洗、转换和分块时,首先需明确相关约定:所需输入、成功信号以及部分失败时的处理方式。这样的清单能确保后续的代码修改保持一致性。 将这一阶段视为输入与经过验证的输出之间的契约。为生成的成果命名,定义成功判定标准,并杜绝无声的半完成状态。
快速纯文本提取
将快速纯文本提取视为可度量的对象时,其效果最佳。在扩大应用范围之前,先记录一个成功的示例、一个失败案例以及回滚说明。在功能结果旁同时记录处理时间以及令牌或查询成本。提前了解成本情况,可避免在从演示环境过渡到共享环境时出现意外费用。应将分块策略与检索策略分开,当质量指标发生变化时,修改其中一项不应迫使重新编写另一项。
from pyspark.sql import functions as F
# Convert VARIANT to JSON string, then extract text content
safe_json_col = F.coalesce(
F.to_json(F.col("parsed_content")),
F.col("parsed_content").cast("string")
)
plain_text_df = parsed_df.withColumn(
"plain_text",
extract_contents_udf()(safe_json_col) # Custom UDF to join text elements
)
使用 LangChain 进行分块
当将 LangChain 的分块功能视为可度量的对象时,其效果最佳。在扩大范围之前,先记录一个理想的转录样本、一个失败案例以及回滚说明。 配置应置于应用程序代码之外。环境文件、密钥存储和功能开关应集中存放,以便操作人员无需查看整个系统结构即可进行审计。 将分块策略与检索策略分开。当质量指标发生变化时,修改其中一项不应强制要求重新编写另一项。
from langchain_text_splitters import RecursiveCharacterTextSplitter
from pyspark.sql.types import StructType, StructField, StringType
import pandas as pd
CHUNK_SIZE = 2000
CHUNK_OVERLAP = 200
splitter = RecursiveCharacterTextSplitter(
chunk_size=CHUNK_SIZE,
chunk_overlap=CHUNK_OVERLAP,
separators=["\n== page ==\n", "== page ==", "\n\n", "\n", " ", ""]
)
schema = StructType([
StructField("path", StringType(), True),
StructField("chunk", StringType(), True),
])
def split_rows(iterator):
for pdf in iterator:
out = []
for _, row in pdf.iterrows():
path, text = row["path"], row["plain_text"]
if isinstance(text, str) and text.strip():
for c in splitter.split_text(text):
if c and c.strip():
out.append((path, c))
yield pd.DataFrame(out, columns=["path", "chunk"])
df_chunks = (
plain_text_df.select("path", "plain_text")
.mapInPandas(split_rows, schema=schema)
)
# Add unique IDs and save
df_chunks = df_chunks.withColumn("id", F.monotonically_increasing_id())
chunked_table = "<YOUR_CATALOG>.<YOUR_SCHEMA>.docs_chunked"
df_chunks.write.format("delta") \
.mode("overwrite") \
.option("mergeSchema", "true") \
.saveAsTable(chunked_table)
步骤 3:构建向量搜索
第三步:将向量搜索视为可度量的模型来构建时效果最佳。在扩大范围之前,先记录一个成功的案例、一个失败案例以及回滚说明。同时记录正常流程和恢复流程。重试机制、人工审核环节以及死信处理都是产品本身的组成部分,而非后续需要补充的功能。应将分块策略与检索策略分开,当质量指标发生变化时,修改其中一项不应迫使重新编写另一项。
启用变更数据源
将“启用变更数据馈送”视为可度量的指标来使用效果最佳。在扩大范围之前,先记录一份完美的操作日志、一个故障案例以及回滚说明。 优先选择小型且可测试的单元,而非庞大的脚本。当某个步骤出现故障时,故障应指向单一责任主体,而非复杂的流程链。 将分块策略与检索策略分开。当质量指标发生变化时,修改其中一项不应迫使重新编写另一项。
ALTER TABLE <YOUR_CATALOG>.<YOUR_SCHEMA>.docs_chunked
SET TBLPROPERTIES (delta.enableChangeDataFeed = true);
创建增量同步索引
将“创建Delta同步索引”视为可测量的对象来处理效果最佳。在扩大范围之前,先记录一份完美的转录文本、一个故障案例以及回滚说明。 把这一阶段视为输入与已验证输出之间的契约。为相关成果命名,明确成功标准,绝不允许出现无声无息的半完成状态。 将分块策略与检索策略分开。当质量指标发生变化时,修改其中一项不应强制要求重新编写另一项。
from databricks.vector_search.client import VectorSearchClient
vsc = VectorSearchClient(disable_notice=True)
index_name = "<YOUR_CATALOG>.<YOUR_SCHEMA>.docs_chunked_index"
vsc.create_delta_sync_index_and_wait(
endpoint_name="<YOUR_VS_ENDPOINT>",
index_name=index_name,
source_table_name="<YOUR_CATALOG>.<YOUR_SCHEMA>.docs_chunked",
primary_key="id",
embedding_source_column="chunk",
embedding_model_endpoint_name="databricks-gte-large-en",
pipeline_type="TRIGGERED",
)
测试检索
将检索测试视为可度量的指标时,其效果最佳。在扩大范围之前,先记录一个成功的案例、一个失败案例以及回滚说明。在功能结果旁同时记录处理时间以及令牌或查询成本。提前了解成本情况,可避免在系统从演示环境过渡到共享环境时出现意外费用。应将分块策略与检索策略分开,当质量指标发生变化时,修改其中一项不应迫使重新编写另一项。
index = vsc.get_index(index_name=index_name)
results = index.similarity_search(
query_text="How does the system prevent overheating?",
columns=["path", "chunk"],
num_results=5,
)
display(results)
第4步:为对话记忆设置Lakebase
第4步:配置Lakebase以用于对话记忆功能。将该功能视为可度量的数据表面时,效果最佳。在扩大应用范围之前,需记录一份优秀的对话转录文本、一个失败案例以及回滚说明。 将配置信息与应用程序代码分开。环境文件、密钥存储和功能开关应集中存放于一处,以便操作人员无需查看整个系统结构即可进行审计。 将分块策略与检索策略分开。当质量指标发生变化时,修改其中一项不应强制要求重新编写另一项的策略。
创建Lakebase自动扩缩容项目
将 Lakebase 自动扩缩容项目视为可度量的对象来管理效果最佳。在扩大范围之前,需记录一份理想运行案例、一个故障场景以及回滚说明。 同时文档化正常流程与恢复流程。重试机制、人工审核环节以及死信处理都是产品本身的组成部分,而非后续需要补充的内容。 应将分块策略与检索策略分开。当质量指标发生变化时,修改其中一项不应强制要求重新编写另一项。 将 Lakebase 自动扩缩容项目视为可度量的对象来管理效果最佳。在扩大范围之前,需记录一份理想运行案例、一个故障场景以及回滚说明。 应将此阶段视为输入与经过验证的输出之间的契约。为相关成果命名,明确成功标准,绝不允许出现无声无息的半完成状态。
通过编程方式获取连接详情
若要通过编程方式获取连接详情,应在修改代码之前明确输入参数、该步骤的负责人以及终止条件。操作人员应能够从已知的检查点重新运行该步骤,而无需猜测隐藏状态。除了功能结果外,还需记录执行时间以及令牌或查询成本。提前了解成本情况可避免在从演示环境切换到共享环境时出现意外费用。必须注明实际作为答案依据的段落;没有引用的话,操作人员就无法区分是幻觉内容还是索引缺失导致的错误。
from databricks.sdk import WorkspaceClient
w = WorkspaceClient()
project_id = "<YOUR_PROJECT_NAME>"
# Get branch and endpoint
branches = list(w.postgres.list_branches(parent=f"projects/{project_id}"))
branch_name = branches[0].name
endpoints = list(w.postgres.list_endpoints(parent=branch_name))
ep = endpoints[0]
HOST = ep.status.hosts.host
ENDPOINT = ep.name
USERNAME = "<YOUR_DATABRICKS_EMAIL>"
print(f"Host: {HOST}")
print(f"Endpoint: {ENDPOINT}")
测试连接
为测试连接,在修改代码之前需先确定输入参数、该步骤的负责人以及结束标准。操作人员应能够从已知的检查点重新运行该步骤,而无需猜测隐藏状态。 配置信息应置于应用程序代码之外。环境文件、密钥存储以及功能标志应集中存放于一个位置,这样操作人员无需查看整个系统结构即可进行审核。 需明确标注支撑答案的具体内容。若没有引用依据,操作人员就无法区分是虚假信息还是索引缺失所致。
import psycopg2
cred = w.postgres.generate_database_credential(endpoint=ENDPOINT)
conn = psycopg2.connect(
host=HOST,
dbname="databricks_postgres",
user=USERNAME,
password=cred.token,
port=5432,
sslmode="require"
)
with conn.cursor() as cur:
cur.execute("SELECT version()")
print(cur.fetchone()[0])
conn.close()
print("✅ Connected to Lakebase!")
创建检查点表
在创建检查点表时,应在修改代码之前明确输入参数、该步骤的负责人以及终止标准。操作人员应能够从已知的检查点重新运行该步骤,而无需猜测隐藏状态。 需同时记录正常流程和故障恢复流程。重试机制、人工审核环节以及错误处理都是产品本身的组成部分,而非后续需要补充的内容。 必须引用实际作为答案依据的段落。如果没有引用,操作人员就无法区分是幻觉内容还是索引缺失导致的错误。
from urllib.parse import quote
from langgraph.checkpoint.postgres import PostgresSaver
cred = w.postgres.generate_database_credential(endpoint=ENDPOINT)
DB_URI = (
f"postgresql://{quote(USERNAME, safe='')}:{quote(cred.token, safe='')}"
f"@{HOST}:5432/databricks_postgres"
f"?sslmode=require"
)
with PostgresSaver.from_conn_string(DB_URI) as checkpointer:
checkpointer.setup()
print("✅ Checkpoint tables created!")
在创建检查点表时,应在修改代码之前明确输入参数、该步骤的负责人以及终止标准。操作人员应能够从已知的检查点重新运行该步骤,而无需猜测隐藏状态。 应将此阶段视为输入参数与经过验证的输出结果之间的契约。为相关成果命名,明确成功判定标准,并杜绝默许部分完成的情况。
第5步:构建上下文感知型智能体
在执行第5步“构建上下文感知型智能体”时,首先列出相关规范:所需输入、成功标志以及部分失败时的处理方式。这样的清单能确保后续的代码修改保持一致性。 在功能结果旁记录执行时间以及令牌或查询成本。提前了解成本情况,可避免在从演示环境过渡到共享环境时出现意外费用。 在调整提示词之前,先使用固定的问题集测试召回率。仅仅更换提示词往往无法改善较差的检索效果。
交互式智能体(笔记本)
在处理交互式智能体(笔记本)时,首先写下相关规范:所需的输入参数、成功信号以及部分失败时的处理方式。这样的检查清单能确保后续的代码修改保持一致性。 将配置信息置于应用程序代码之外。环境文件、密钥存储以及功能开关应集中存放,这样操作人员无需查看整个系统结构即可进行审计。 在调整提示词之前,先使用固定的问题集来测试召回率。仅仅更换提示词往往无法解决检索效果不佳的问题。
from urllib.parse import quote
from langchain.agents import create_agent
from databricks_langchain import ChatDatabricks, VectorSearchRetrieverTool
from langgraph.checkpoint.postgres import PostgresSaver
from databricks.sdk import WorkspaceClient
import psycopg
from psycopg.rows import dict_row
def get_lakebase_checkpointer(host: str, endpoint: str, username: str):
"""Create a PostgresSaver backed by Lakebase Autoscaling."""
w = WorkspaceClient()
cred = w.postgres.generate_database_credential(endpoint=endpoint)
db_uri = (
f"postgresql://{quote(username, safe='')}:{quote(cred.token, safe='')}"
f"@{host}:5432/databricks_postgres"
f"?sslmode=require"
)
# IMPORTANT: Use psycopg.connect directly, not from_conn_string
# from_conn_string returns a context manager, not a persistent instance
conn = psycopg.connect(db_uri, autocommit=True, row_factory=dict_row)
checkpointer = PostgresSaver(conn=conn)
checkpointer.setup()
return checkpointer
def build_agent(llm_endpoint: str, index_name: str, num_results: int = 3):
model = ChatDatabricks(endpoint=llm_endpoint, max_tokens=500)
vs_tool = VectorSearchRetrieverTool(
name="knowledge_search",
index_name=index_name,
description="Search knowledge base for relevant information",
num_results=num_results,
)
# Lakebase-backed checkpointer instead of InMemorySaver
checkpointer = get_lakebase_checkpointer(HOST, ENDPOINT, USERNAME)
system_prompt = """You are a Knowledge Assistant. Respond in a clear,
professional tone. Use only verified information from the provided documents.
If the answer cannot be found, clearly state that."""
return create_agent(
model=model,
tools=[vs_tool],
system_prompt=system_prompt,
checkpointer=checkpointer,
)
测试多轮对话功能
在处理多轮对话测试时,首先需明确相关约定:所需的输入参数、成功标志以及部分失败时的处理方式。这样的清单能确保后续的代码修改始终符合预期。 同时记录正常流程与异常恢复路径。重试机制、人工审核环节以及错误消息处理都是产品功能的一部分,而非后续需要补充的内容。 在调整提示词之前,先使用固定的问题集来评估信息检索的准确率。仅仅更换提示词往往无法解决检索效果不佳的问题。
agent = build_agent("<YOUR_LLM_ENDPOINT>", "<YOUR_INDEX_NAME>", 3)
# STABLE thread_id - this is what enables context awareness
config = {"configurable": {"thread_id": "demo-session-001"}}
# Turn 1
r1 = agent.invoke(
{"messages": [{"role": "user", "content": "What is the Orion system?"}]},
config=config
)
print("Turn 1:", r1['messages'][-1].content)
# Turn 2 - agent should know "it" = Orion
r2 = agent.invoke(
{"messages": [{"role": "user", "content": "How does it handle overheating?"}]},
config=config
)
print("Turn 2:", r2['messages'][-1].content)
在处理多轮对话测试时,首先需明确相关约定:所需的输入参数、成功标志以及部分失败时的处理方式。这样的清单能确保后续的代码修改始终符合预期。 应将此阶段视为输入与验证后输出之间的契约。为相关产出命名,明确成功判定标准,杜绝默许部分完成的情况。
第6步:生产环境代理代码(agent.py)
第6步:将生产环境代理代码(agent.py)视为可度量的对象来处理效果最佳。在扩大范围之前,先记录一个理想案例、一个失败案例以及回滚说明。在功能结果旁同时记录执行时间以及令牌或查询成本。提前了解成本情况,可避免在从演示环境过渡到共享环境时出现意外费用。应将分块策略与检索策略分开,当质量指标发生变化时,修改其中一项无需强制重写另一项。
# agent.py
import os
from uuid import uuid4
from typing import Any, Dict, List
from urllib.parse import quote
import yaml
import mlflow
import psycopg
from psycopg.rows import dict_row
from mlflow.pyfunc import ResponsesAgent
from mlflow.types.responses import ResponsesAgentRequest, ResponsesAgentResponse
from langchain.agents import create_agent
from databricks_langchain import ChatDatabricks, VectorSearchRetrieverTool
from langgraph.checkpoint.postgres import PostgresSaver
from databricks.sdk import WorkspaceClient
def _load_config(path: str = "agent-config.yaml") -> Dict[str, Any]:
if not os.path.exists(path):
raise FileNotFoundError(f"Config file not found at '{path}'")
with open(path, "r", encoding="utf-8") as f:
cfg = yaml.safe_load(f) or {}
llm_endpoint = cfg.get("llm_endpoint_name")
vs = cfg.get("vector_search", {}) or {}
index_name = vs.get("index_name")
num_results = int(vs.get("num_results", 3))
lakebase = cfg.get("lakebase", {}) or {}
return {
"llm_endpoint_name": llm_endpoint,
"vs_index_name": index_name,
"vs_num_results": num_results,
"lakebase_host": lakebase.get("host"),
"lakebase_endpoint": lakebase.get("endpoint"),
"lakebase_user": lakebase.get("user"),
}
def get_lakebase_checkpointer(host, endpoint, user):
w = WorkspaceClient()
cred = w.postgres.generate_database_credential(endpoint=endpoint)
db_uri = (
f"postgresql://{quote(user, safe='')}:{quote(cred.token, safe='')}"
f"@{host}:5432/databricks_postgres?sslmode=require"
)
conn = psycopg.connect(db_uri, autocommit=True, row_factory=dict_row)
checkpointer = PostgresSaver(conn=conn)
checkpointer.setup()
return checkpointer
def build_agent(llm_endpoint, index_name, num_results,
lakebase_host, lakebase_endpoint, lakebase_user):
model = ChatDatabricks(endpoint=llm_endpoint, max_tokens=500)
vs_tool = VectorSearchRetrieverTool(
name="knowledge_search",
index_name=index_name,
description="Search knowledge base for relevant information",
num_results=num_results,
)
checkpointer = get_lakebase_checkpointer(
lakebase_host, lakebase_endpoint, lakebase_user
)
system_prompt = (
"You are a Knowledge Assistant. Respond in a clear, professional tone. "
"Use only verified information from the provided documents. "
"If the answer cannot be found, clearly state that."
)
return create_agent(
model=model, tools=[vs_tool],
system_prompt=system_prompt, checkpointer=checkpointer,
)
def _last_user_text(messages):
user_msgs = [m for m in messages if m.get("role") == "user"]
return str(user_msgs[-1].get("content", "")) if user_msgs else ""
class LangChainResponsesAgent(ResponsesAgent):
def __init__(self):
cfg = _load_config()
self._agent = build_agent(
cfg["llm_endpoint_name"], cfg["vs_index_name"],
cfg["vs_num_results"], cfg["lakebase_host"],
cfg["lakebase_endpoint"], cfg["lakebase_user"],
)
def predict(self, request: ResponsesAgentRequest) -> ResponsesAgentResponse:
msgs = [m.model_dump() for m in request.input]
custom_inputs = dict(request.custom_inputs or {})
thread_id = custom_inputs.get("thread_id", f"session-{uuid4()}")
result = self._agent.invoke(
{"messages": msgs},
config={"configurable": {"thread_id": thread_id}},
)
try:
text = result["messages"][-1].content
except Exception:
text = str(result)
return ResponsesAgentResponse(
output=[self.create_text_output_item(text, str(uuid4()))],
custom_outputs={"thread_id": thread_id},
)
AGENT = LangChainResponsesAgent()
mlflow.models.set_model(AGENT)
配置文件(agent-config.yaml)
将配置文件(agent-config.yaml)视为可度量的对象来管理效果最佳。在扩大应用范围之前,先记录一份理想的运行日志、一个故障案例以及回滚说明。
应将配置文件与应用程序代码分开存放。环境配置文件、密钥存储和功能开关应集中于一个位置,以便操作人员无需查看整个系统结构即可进行审计。
需将分块策略与检索策略分开处理。当质量指标发生变化时,修改其中一项不应强制要求重新编写另一项。
llm_endpoint_name: <YOUR_LLM_ENDPOINT>
vector_search:
index_name: <YOUR_CATALOG>.<YOUR_SCHEMA>.docs_chunked_index
num_results: 3
lakebase:
host: <YOUR_LAKEBASE_HOST>
endpoint: projects/<YOUR_PROJECT>/branches/production/endpoints/primary
user: <YOUR_DATABRICKS_EMAIL>
第7步:记录、注册与部署
将“第7步:记录、注册与部署”视为可度量的对象来管理效果最佳。在扩大应用范围之前,先记录一份理想的运行日志、一个故障案例以及回滚说明。
将数据记录到MLflow
import mlflow
from importlib.metadata import version as get_version
from mlflow.models.resources import DatabricksVectorSearchIndex, DatabricksServingEndpoint
resources = [
DatabricksVectorSearchIndex(index_name="<YOUR_CATALOG>.<YOUR_SCHEMA>.docs_chunked_index"),
DatabricksServingEndpoint(endpoint_name="<YOUR_LLM_ENDPOINT>"),
]
with mlflow.start_run():
mlflow.set_tags({
"model_type": "retrieval_agent",
"framework": "langchain",
"memory": "lakebase_autoscaling",
})
logged_agent_info = mlflow.pyfunc.log_model(
name="knowledge_assistant",
python_model="agent.py",
code_paths=["agent-config.yaml"],
input_example={"input": [{"role": "user", "content": "What is Orion?"}]},
pip_requirements=[
f"databricks-vectorsearch=={get_version('databricks-vectorsearch')}",
f"databricks-langchain=={get_version('databricks-langchain')}",
f"langchain=={get_version('langchain')}",
f"mlflow=={get_version('mlflow')}",
"langgraph-checkpoint-postgres",
"psycopg[binary]",
"databricks-sdk>=0.89.0",
],
resources=resources,
)
model_uri = logged_agent_info.model_uri
在Unity Catalog中注册
mlflow.set_registry_uri("databricks-uc")
UC_MODEL_NAME = "<YOUR_CATALOG>.<YOUR_SCHEMA>.knowledge_assistant"
uc_info = mlflow.register_model(model_uri=model_uri, name=UC_MODEL_NAME)
print(f"✅ Registered: {UC_MODEL_NAME} v{uc_info.version}")
部署
from databricks import agents
deployment = agents.deploy(
model_name=UC_MODEL_NAME,
model_version=uc_info.version,
scale_to_zero_enabled=True,
)
print(f"✅ Endpoint: {deployment.query_endpoint}")
最终结果
ws = WorkspaceClient()
client = ws.serving_endpoints.get_open_ai_client()
session = "user-session-042"
# Turn 1
r1 = client.responses.create(
model="knowledge_assistant",
input=[{"role": "user", "content": "What is the Orion motion controller?"}],
extra_body={"custom_inputs": {"thread_id": session}}
)
# Turn 2 - "it" resolves correctly to Orion
r2 = client.responses.create(
model="knowledge_assistant",
input=[{"role": "user", "content": "How does it prevent overheating?"}],
extra_body={"custom_inputs": {"thread_id": session}}
)
验证内存持久性的快速检查:
import json
from uuid import uuid4
thread_id = f"memory-test-{uuid4()}"
# --- 1. Use custom input instead of input_example ---
custom_input = {
"input": [{"role": "user", "content": "What are the main components of Orion?"}],
"custom_inputs": {"thread_id": thread_id},
}
print("=" * 50)
print("TEST 1: Custom input (no input_example needed)")
print("=" * 50)
# Use output_path to capture results (without it, mlflow.models.predict returns None)
mlflow.models.predict(
model_uri=model_uri,
input_data=custom_input,
env_manager="uv",
output_path="/tmp/result_1.json",
)
with open("/tmp/result_1.json", "r") as f:
result_1 = json.load(f)
thread_id_1 = result_1["custom_outputs"]["thread_id"]
response_1 = result_1["output"][0]["content"][0]["text"]
print(f"Thread ID: {thread_id_1}")
print(f"Response: {response_1[:300]}...")
# --- 2. Test memory persistence with same thread_id ---
follow_up_input = {
"input": [
{
"role": "user",
"content": "Can you elaborate more on the first component you mentioned?",
}
],
"custom_inputs": {"thread_id": thread_id}, # Reuse same thread
}
print("\n" + "=" * 50)
print("TEST 2: Follow-up on same thread (memory test)")
print("=" * 50)
mlflow.models.predict(
model_uri=model_uri,
input_data=follow_up_input,
env_manager="uv",
output_path="/tmp/result_2.json",
)
with open("/tmp/result_2.json", "r") as f:
result_2 = json.load(f)
thread_id_2 = result_2["custom_outputs"]["thread_id"]
response_2 = result_2["output"][0]["content"][0]["text"]
print(f"Thread ID: {thread_id_2}")
print(f"Response: {response_2[:300]}...")
# --- 3. Verify memory with actual conditions ---
print("\n" + "=" * 50)
print("MEMORY CHECK")
print("=" * 50)
# Check 1: Thread IDs match
if thread_id_1 == thread_id_2:
print(f"✅ Thread ID match: {thread_id_1}")
else:
print(f"❌ Thread ID mismatch! Call 1: {thread_id_1}, Call 2: {thread_id_2}")
# Check 2: Follow-up response references context from the first response
follow_up_lower = response_2.lower()
if len(response_2) > 50 and any(
keyword in follow_up_lower
for keyword in [
"motion",
"vision",
"cognition",
"communication",
"subsystem",
"component",
]
):
print(
"✅ Follow-up response references components from the first answer — memory is intact!"
)
else:
print(
"⚠️ Follow-up response may not reference the first answer. Manual review recommended."
)
print(f" Follow-up preview: {response_2[:200]}")
print(
"\n✅ Lakebase Postgres checkpointing is working correctly!"
if thread_id_1 == thread_id_2
else "\n❌ Memory persistence test FAILED."
)