Improve context handling for workflow runs by auto-detecting target URLs and setting the previous assistant message as the default target ID.

This commit is contained in:
2026-06-15 15:46:25 +02:00
parent f0edcbccd5
commit c7eed48fce

View File

@@ -343,6 +343,12 @@ async def create_workflow_run(request: WorkflowRunRequest):
session = None session = None
user_message = None user_message = None
run_attachments = request.attachments run_attachments = request.attachments
sample_inputs = dict(request.sample_inputs)
target_message_id = sample_inputs.get("target_message_id") or sample_inputs.get("message_id")
target_url = sample_inputs.get("target_url") or sample_inputs.get("url")
if not target_url:
url_match = re.search(r"https?://[^\s<>'\"]+", request.message or "")
target_url = url_match.group(0).rstrip(".,);]") if url_match else None
if request.session_id: if request.session_id:
session = db.query(chat_models.ChatSession).filter(chat_models.ChatSession.session_id == request.session_id).first() session = db.query(chat_models.ChatSession).filter(chat_models.ChatSession.session_id == request.session_id).first()
if session is None and request.regenerate_index is not None: if session is None and request.regenerate_index is not None:
@@ -351,6 +357,13 @@ async def create_workflow_run(request: WorkflowRunRequest):
session = chat_models.ChatSession(session_id=request.session_id) session = chat_models.ChatSession(session_id=request.session_id)
db.add(session) db.add(session)
db.flush() db.flush()
if not target_message_id:
previous_assistant = db.query(chat_models.ChatMessage).filter(
chat_models.ChatMessage.session_pk == session.id,
chat_models.ChatMessage.role == "assistant",
).order_by(chat_models.ChatMessage.created_at.desc(), chat_models.ChatMessage.id.desc()).first()
if previous_assistant:
target_message_id = previous_assistant.message_id
from .. import main as main_module from .. import main as main_module
if request.regenerate_index is not None: if request.regenerate_index is not None:
rows = db.query(chat_models.ChatMessage).filter( rows = db.query(chat_models.ChatMessage).filter(
@@ -386,7 +399,6 @@ async def create_workflow_run(request: WorkflowRunRequest):
messages = [{"role": row.role, "content": row.content} for row in rows] messages = [{"role": row.role, "content": row.content} for row in rows]
else: else:
messages = [{"role": "user", "content": request.message}] messages = [{"role": "user", "content": request.message}]
sample_inputs = dict(request.sample_inputs)
sample_inputs.setdefault("prompt", request.message) sample_inputs.setdefault("prompt", request.message)
sample_inputs.setdefault("session_id", request.session_id) sample_inputs.setdefault("session_id", request.session_id)
run_values = { run_values = {
@@ -405,8 +417,8 @@ async def create_workflow_run(request: WorkflowRunRequest):
"manual_library_enabled": bool(request.library_slug), "manual_library_enabled": bool(request.library_slug),
"web_search_enabled": request.web_search_enabled, "web_search_enabled": request.web_search_enabled,
"show_thinking": request.show_thinking, "show_thinking": request.show_thinking,
"target_message_id": sample_inputs.get("target_message_id") or sample_inputs.get("message_id"), "target_message_id": target_message_id,
"target_url": sample_inputs.get("target_url") or sample_inputs.get("url"), "target_url": target_url,
} }
run = WorkflowRun( run = WorkflowRun(
workflow_id=workflow.id, workflow_revision_id=revision.id, session_id=request.session_id, workflow_id=workflow.id, workflow_revision_id=revision.id, session_id=request.session_id,