mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-09-05 17:17:14 +00:00
Compare commits
34 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 59538fd5cd | |||
| 235b33c335 | |||
| c588f9dfc1 | |||
| d770de7f7f | |||
| 417cca016c | |||
| cb0cc4d06a | |||
| cb0bb44db8 | |||
| 6609bebeec | |||
| 192b69b0fb | |||
| 543fbd8ca0 | |||
| 804448b6cd | |||
| 5d4e40459f | |||
| b5c43cc113 | |||
| 8ebfcd963a | |||
| 127198675e | |||
| 44fb188994 | |||
| 265385a563 | |||
| 253cc6cbea | |||
| f99d3022e8 | |||
| 5c5614667a | |||
| 313d553271 | |||
| bb7db53447 | |||
| 27c0d344bf | |||
| c088dc114f | |||
| 37c74b0622 | |||
| c9f7911efe | |||
| 75fdfe6806 | |||
| eb9f38b102 | |||
| fc40d3c949 | |||
| d176a448e0 | |||
| ada4c30f85 | |||
| 32c9eaff45 | |||
| 9706ee2d53 | |||
| e7c9bc69d3 |
@@ -51,7 +51,7 @@ LangBot is an **open-source, production-grade platform** for building AI-powered
|
||||
|
||||
[→ Learn more about all features](https://link.langbot.app/en/docs/features)
|
||||
|
||||
📍 Practical guides: [deploy a multi-platform AI bot in 5 minutes](https://langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [connect DeepSeek to WeChat, Discord, and Telegram](https://langbot.app/en/blog/connect-deepseek-to-wechat/), [run a Dify Agent in Discord, Telegram, and Slack](https://langbot.app/en/blog/dify-agent-discord-telegram-slack/), and [build an n8n-powered chatbot](https://langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
📍 Practical guides: [deploy a multi-platform AI bot in 5 minutes](https://blog.langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [connect DeepSeek to WeChat, Discord, and Telegram](https://blog.langbot.app/en/blog/connect-deepseek-to-wechat/), [run a Dify Agent in Discord, Telegram, and Slack](https://blog.langbot.app/en/blog/dify-agent-discord-telegram-slack/), and [build an n8n-powered chatbot](https://blog.langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
|
||||
---
|
||||
|
||||
|
||||
+1
-1
@@ -51,7 +51,7 @@ LangBot 是一个**开源的生产级平台**,用于构建 AI 驱动的即时
|
||||
|
||||
[→ 了解更多功能特性](https://link.langbot.app/zh/docs/features)
|
||||
|
||||
📍 实践指南:[5 分钟部署多平台 AI 机器人](https://langbot.app/zh/blog/deploy-ai-bot-in-5-minutes/)、[将 DeepSeek 接入微信、企业微信与 Discord](https://langbot.app/zh/blog/connect-deepseek-to-wechat/)、[让 Dify Agent 跑在 Discord、Telegram 和 Slack 上](https://langbot.app/zh/blog/dify-agent-discord-telegram-slack/),以及[用 n8n 构建多平台 AI 聊天机器人](https://langbot.app/zh/blog/n8n-multi-platform-ai-chatbot/)。
|
||||
📍 实践指南:[5 分钟部署多平台 AI 机器人](https://blog.langbot.app/zh/blog/deploy-ai-bot-in-5-minutes/)、[将 DeepSeek 接入微信、企业微信与 Discord](https://blog.langbot.app/zh/blog/connect-deepseek-to-wechat/)、[让 Dify Agent 跑在 Discord、Telegram 和 Slack 上](https://blog.langbot.app/zh/blog/dify-agent-discord-telegram-slack/),以及[用 n8n 构建多平台 AI 聊天机器人](https://blog.langbot.app/zh/blog/n8n-multi-platform-ai-chatbot/)。
|
||||
|
||||
---
|
||||
|
||||
|
||||
+1
-1
@@ -50,7 +50,7 @@ LangBot es una **plataforma de código abierto y grado de producción** para con
|
||||
|
||||
[→ Conocer más sobre todas las funcionalidades](https://link.langbot.app/en/docs/features)
|
||||
|
||||
📍 Guías prácticas: [desplegar un bot de IA multiplataforma en 5 minutos](https://langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [conectar DeepSeek a WeChat, Discord y Telegram](https://langbot.app/en/blog/connect-deepseek-to-wechat/), [ejecutar un Dify Agent en Discord, Telegram y Slack](https://langbot.app/en/blog/dify-agent-discord-telegram-slack/) y [crear un chatbot con n8n](https://langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
📍 Guías prácticas: [desplegar un bot de IA multiplataforma en 5 minutos](https://blog.langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [conectar DeepSeek a WeChat, Discord y Telegram](https://blog.langbot.app/en/blog/connect-deepseek-to-wechat/), [ejecutar un Dify Agent en Discord, Telegram y Slack](https://blog.langbot.app/en/blog/dify-agent-discord-telegram-slack/) y [crear un chatbot con n8n](https://blog.langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
|
||||
---
|
||||
|
||||
|
||||
+1
-1
@@ -50,7 +50,7 @@ LangBot est une **plateforme open-source de niveau production** pour créer des
|
||||
|
||||
[→ En savoir plus sur toutes les fonctionnalités](https://link.langbot.app/en/docs/features)
|
||||
|
||||
📍 Guides pratiques : [déployer un bot IA multiplateforme en 5 minutes](https://langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [connecter DeepSeek à WeChat, Discord et Telegram](https://langbot.app/en/blog/connect-deepseek-to-wechat/), [exécuter un Dify Agent dans Discord, Telegram et Slack](https://langbot.app/en/blog/dify-agent-discord-telegram-slack/) et [créer un chatbot avec n8n](https://langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
📍 Guides pratiques : [déployer un bot IA multiplateforme en 5 minutes](https://blog.langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [connecter DeepSeek à WeChat, Discord et Telegram](https://blog.langbot.app/en/blog/connect-deepseek-to-wechat/), [exécuter un Dify Agent dans Discord, Telegram et Slack](https://blog.langbot.app/en/blog/dify-agent-discord-telegram-slack/) et [créer un chatbot avec n8n](https://blog.langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
|
||||
---
|
||||
|
||||
|
||||
+1
-1
@@ -50,7 +50,7 @@ LangBot は、AI搭載のインスタントメッセージングボットを構
|
||||
|
||||
[→ すべての機能について詳しく見る](https://link.langbot.app/ja/docs/features)
|
||||
|
||||
📍 実践ガイド: [5分でマルチプラットフォームAIボットをデプロイ](https://langbot.app/en/blog/deploy-ai-bot-in-5-minutes/)、[DeepSeekをWeChat・Discord・Telegramに接続](https://langbot.app/en/blog/connect-deepseek-to-wechat/)、[Dify AgentをDiscord・Telegram・Slackで動かす](https://langbot.app/en/blog/dify-agent-discord-telegram-slack/)、[n8n連携チャットボットを構築](https://langbot.app/en/blog/n8n-multi-platform-ai-chatbot/)。
|
||||
📍 実践ガイド: [5分でマルチプラットフォームAIボットをデプロイ](https://blog.langbot.app/en/blog/deploy-ai-bot-in-5-minutes/)、[DeepSeekをWeChat・Discord・Telegramに接続](https://blog.langbot.app/en/blog/connect-deepseek-to-wechat/)、[Dify AgentをDiscord・Telegram・Slackで動かす](https://blog.langbot.app/en/blog/dify-agent-discord-telegram-slack/)、[n8n連携チャットボットを構築](https://blog.langbot.app/en/blog/n8n-multi-platform-ai-chatbot/)。
|
||||
|
||||
---
|
||||
|
||||
|
||||
+1
-1
@@ -50,7 +50,7 @@ LangBot은 AI 기반 인스턴트 메시징 봇을 구축하기 위한 **오픈
|
||||
|
||||
[→ 모든 기능 자세히 보기](https://link.langbot.app/en/docs/features)
|
||||
|
||||
📍 실전 가이드: [5분 만에 멀티 플랫폼 AI 봇 배포하기](https://langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [DeepSeek를 WeChat, Discord, Telegram에 연결하기](https://langbot.app/en/blog/connect-deepseek-to-wechat/), [Dify Agent를 Discord, Telegram, Slack에서 실행하기](https://langbot.app/en/blog/dify-agent-discord-telegram-slack/), [n8n 기반 챗봇 만들기](https://langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
📍 실전 가이드: [5분 만에 멀티 플랫폼 AI 봇 배포하기](https://blog.langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [DeepSeek를 WeChat, Discord, Telegram에 연결하기](https://blog.langbot.app/en/blog/connect-deepseek-to-wechat/), [Dify Agent를 Discord, Telegram, Slack에서 실행하기](https://blog.langbot.app/en/blog/dify-agent-discord-telegram-slack/), [n8n 기반 챗봇 만들기](https://blog.langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
|
||||
---
|
||||
|
||||
|
||||
+1
-1
@@ -50,7 +50,7 @@ LangBot — это **платформа с открытым исходным к
|
||||
|
||||
[→ Подробнее обо всех возможностях](https://link.langbot.app/en/docs/features)
|
||||
|
||||
📍 Практические руководства: [развернуть мультиплатформенного ИИ-бота за 5 минут](https://langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [подключить DeepSeek к WeChat, Discord и Telegram](https://langbot.app/en/blog/connect-deepseek-to-wechat/), [запустить Dify Agent в Discord, Telegram и Slack](https://langbot.app/en/blog/dify-agent-discord-telegram-slack/) и [создать чат-бота на n8n](https://langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
📍 Практические руководства: [развернуть мультиплатформенного ИИ-бота за 5 минут](https://blog.langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [подключить DeepSeek к WeChat, Discord и Telegram](https://blog.langbot.app/en/blog/connect-deepseek-to-wechat/), [запустить Dify Agent в Discord, Telegram и Slack](https://blog.langbot.app/en/blog/dify-agent-discord-telegram-slack/) и [создать чат-бота на n8n](https://blog.langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
|
||||
---
|
||||
|
||||
|
||||
+1
-1
@@ -52,7 +52,7 @@ LangBot 是一個**開源的生產級平台**,用於建構 AI 驅動的即時
|
||||
|
||||
[→ 了解更多功能特性](https://link.langbot.app/zh/docs/features)
|
||||
|
||||
📍 實踐指南:[5 分鐘部署多平台 AI 機器人](https://langbot.app/zh/blog/deploy-ai-bot-in-5-minutes/)、[將 DeepSeek 接入微信、企業微信與 Discord](https://langbot.app/zh/blog/connect-deepseek-to-wechat/)、[讓 Dify Agent 跑在 Discord、Telegram 和 Slack 上](https://langbot.app/zh/blog/dify-agent-discord-telegram-slack/),以及[用 n8n 建構多平台 AI 聊天機器人](https://langbot.app/zh/blog/n8n-multi-platform-ai-chatbot/)。
|
||||
📍 實踐指南:[5 分鐘部署多平台 AI 機器人](https://blog.langbot.app/zh/blog/deploy-ai-bot-in-5-minutes/)、[將 DeepSeek 接入微信、企業微信與 Discord](https://blog.langbot.app/zh/blog/connect-deepseek-to-wechat/)、[讓 Dify Agent 跑在 Discord、Telegram 和 Slack 上](https://blog.langbot.app/zh/blog/dify-agent-discord-telegram-slack/),以及[用 n8n 建構多平台 AI 聊天機器人](https://blog.langbot.app/zh/blog/n8n-multi-platform-ai-chatbot/)。
|
||||
|
||||
---
|
||||
|
||||
|
||||
+1
-1
@@ -50,7 +50,7 @@ LangBot là một **nền tảng mã nguồn mở, cấp sản xuất** để x
|
||||
|
||||
[→ Tìm hiểu thêm về tất cả tính năng](https://link.langbot.app/en/docs/features)
|
||||
|
||||
📍 Hướng dẫn thực hành: [triển khai bot AI đa nền tảng trong 5 phút](https://langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [kết nối DeepSeek với WeChat, Discord và Telegram](https://langbot.app/en/blog/connect-deepseek-to-wechat/), [chạy Dify Agent trên Discord, Telegram và Slack](https://langbot.app/en/blog/dify-agent-discord-telegram-slack/) và [xây dựng chatbot với n8n](https://langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
📍 Hướng dẫn thực hành: [triển khai bot AI đa nền tảng trong 5 phút](https://blog.langbot.app/en/blog/deploy-ai-bot-in-5-minutes/), [kết nối DeepSeek với WeChat, Discord và Telegram](https://blog.langbot.app/en/blog/connect-deepseek-to-wechat/), [chạy Dify Agent trên Discord, Telegram và Slack](https://blog.langbot.app/en/blog/dify-agent-discord-telegram-slack/) và [xây dựng chatbot với n8n](https://blog.langbot.app/en/blog/n8n-multi-platform-ai-chatbot/).
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -0,0 +1,163 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Compare YAML node definitions with frontend node-configs."""
|
||||
|
||||
import yaml
|
||||
import os
|
||||
import re
|
||||
import json
|
||||
|
||||
# 1. Parse YAML files
|
||||
yaml_dir = 'src/langbot/templates/metadata/nodes'
|
||||
yaml_nodes = {}
|
||||
|
||||
for filename in sorted(os.listdir(yaml_dir)):
|
||||
if filename.endswith('.yaml'):
|
||||
filepath = os.path.join(yaml_dir, filename)
|
||||
with open(filepath, 'r') as f:
|
||||
data = yaml.safe_load(f)
|
||||
node_name = data.get('name', filename.replace('.yaml', ''))
|
||||
yaml_nodes[node_name] = {
|
||||
'category': data.get('category', ''),
|
||||
'inputs': [i['name'] for i in data.get('inputs', [])],
|
||||
'outputs': [o['name'] for o in data.get('outputs', [])],
|
||||
'config': [c['name'] for c in data.get('config', [])]
|
||||
}
|
||||
|
||||
# 2. Parse frontend node-configs TypeScript files
|
||||
node_configs_dir = 'web/src/app/home/workflows/components/workflow-editor/node-configs'
|
||||
|
||||
frontend_nodes = {}
|
||||
|
||||
def parse_ts_file(filepath):
|
||||
"""Parse a TypeScript file to extract node configurations."""
|
||||
with open(filepath, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
# Find all node type definitions
|
||||
# Pattern: nodeType: 'xxx'
|
||||
node_type_pattern = r"nodeType:\s*'([^']+)'"
|
||||
node_types = re.findall(node_type_pattern, content)
|
||||
|
||||
# For each node type, extract inputs, outputs, and config
|
||||
for node_type in node_types:
|
||||
# Find the config object for this node type
|
||||
# Look for the section between this nodeType and the next one or end of object
|
||||
pattern = rf"nodeType:\s*'({re.escape(node_type)})'.*?(?=nodeType:|export\s+(const|function)|$)"
|
||||
match = re.search(pattern, content, re.DOTALL)
|
||||
|
||||
if match:
|
||||
section = match.group(0)
|
||||
|
||||
# Extract inputs
|
||||
inputs = re.findall(r"createInput\('([^']+)'", section)
|
||||
|
||||
# Extract outputs
|
||||
outputs = re.findall(r"createOutput\('([^']+)'", section)
|
||||
|
||||
# Extract config names
|
||||
config_names = re.findall(r"name:\s*'([^']+)'", section)
|
||||
# Remove duplicates while preserving order
|
||||
seen = set()
|
||||
unique_config = []
|
||||
for c in config_names:
|
||||
if c not in seen:
|
||||
seen.add(c)
|
||||
unique_config.append(c)
|
||||
|
||||
frontend_nodes[node_type] = {
|
||||
'inputs': inputs,
|
||||
'outputs': outputs,
|
||||
'config': unique_config
|
||||
}
|
||||
|
||||
# Parse all config files
|
||||
for filename in os.listdir(node_configs_dir):
|
||||
if filename.endswith('.ts') and filename != 'types.ts' and filename != 'index.ts':
|
||||
filepath = os.path.join(node_configs_dir, filename)
|
||||
parse_ts_file(filepath)
|
||||
|
||||
# 3. Compare and report differences
|
||||
print("=" * 80)
|
||||
print("WORKFLOW NODE COMPARISON REPORT: YAML vs Frontend")
|
||||
print("=" * 80)
|
||||
|
||||
all_node_types = sorted(set(list(yaml_nodes.keys()) + list(frontend_nodes.keys())))
|
||||
|
||||
discrepancies = []
|
||||
|
||||
for node_type in all_node_types:
|
||||
yaml_def = yaml_nodes.get(node_type)
|
||||
frontend_def = frontend_nodes.get(node_type)
|
||||
|
||||
node_discrepancies = []
|
||||
|
||||
if not yaml_def:
|
||||
print(f"\n⚠️ {node_type}: ONLY in frontend (not in YAML)")
|
||||
continue
|
||||
if not frontend_def:
|
||||
print(f"\n⚠️ {node_type}: ONLY in YAML (not in frontend)")
|
||||
continue
|
||||
|
||||
# Compare inputs
|
||||
yaml_inputs = set(yaml_def['inputs'])
|
||||
frontend_inputs = set(frontend_def['inputs'])
|
||||
if yaml_inputs != frontend_inputs:
|
||||
only_yaml = yaml_inputs - frontend_inputs
|
||||
only_frontend = frontend_inputs - yaml_inputs
|
||||
node_discrepancies.append({
|
||||
'type': 'inputs',
|
||||
'only_yaml': list(only_yaml),
|
||||
'only_frontend': list(only_frontend)
|
||||
})
|
||||
|
||||
# Compare outputs
|
||||
yaml_outputs = set(yaml_def['outputs'])
|
||||
frontend_outputs = set(frontend_def['outputs'])
|
||||
if yaml_outputs != frontend_outputs:
|
||||
only_yaml = yaml_outputs - frontend_outputs
|
||||
only_frontend = frontend_outputs - yaml_outputs
|
||||
node_discrepancies.append({
|
||||
'type': 'outputs',
|
||||
'only_yaml': list(only_yaml),
|
||||
'only_frontend': list(only_frontend)
|
||||
})
|
||||
|
||||
# Compare config
|
||||
yaml_config = set(yaml_def['config'])
|
||||
frontend_config = set(frontend_def['config'])
|
||||
if yaml_config != frontend_config:
|
||||
only_yaml = yaml_config - frontend_config
|
||||
only_frontend = frontend_config - yaml_config
|
||||
node_discrepancies.append({
|
||||
'type': 'config',
|
||||
'only_yaml': list(only_yaml),
|
||||
'only_frontend': list(only_frontend)
|
||||
})
|
||||
|
||||
if node_discrepancies:
|
||||
print(f"\n❌ {node_type} ({yaml_def['category']}): HAS DISCREPANCIES")
|
||||
for d in node_discrepancies:
|
||||
print(f" {d['type']}:")
|
||||
if d['only_yaml']:
|
||||
print(f" Only in YAML: {d['only_yaml']}")
|
||||
if d['only_frontend']:
|
||||
print(f" Only in Frontend: {d['only_frontend']}")
|
||||
discrepancies.append((node_type, node_discrepancies))
|
||||
else:
|
||||
print(f"\n✅ {node_type} ({yaml_def['category']}): OK")
|
||||
|
||||
print(f"\n{'=' * 80}")
|
||||
print(f"SUMMARY: {len(discrepancies)} nodes with discrepancies out of {len(all_node_types)} total")
|
||||
print(f"{'=' * 80}")
|
||||
|
||||
# Output as JSON for further processing
|
||||
output = {
|
||||
'yaml_nodes': {k: v for k, v in yaml_nodes.items()},
|
||||
'frontend_nodes': {k: v for k, v in frontend_nodes.items()},
|
||||
'discrepancies': {k: v for k, v in discrepancies}
|
||||
}
|
||||
|
||||
with open('node_comparison.json', 'w') as f:
|
||||
json.dump(output, f, indent=2)
|
||||
|
||||
print(f"\nDetailed comparison saved to node_comparison.json")
|
||||
@@ -1,169 +0,0 @@
|
||||
# Valkey Search Vector Database Integration
|
||||
|
||||
This document describes how to use **Valkey Search** (the search/vector module bundled in
|
||||
`valkey/valkey-bundle`) as the vector database backend for LangBot's knowledge base (RAG)
|
||||
feature.
|
||||
|
||||
## What is Valkey Search?
|
||||
|
||||
**Valkey Search** is a module that adds vector similarity search and full-text search to
|
||||
[Valkey](https://valkey.io/), the open-source, BSD-licensed in-memory data store forked from
|
||||
Redis OSS. It is distributed in the `valkey/valkey-bundle` image alongside other modules
|
||||
(JSON, Bloom, LDAP).
|
||||
|
||||
LangBot talks to Valkey through the official [`valkey-glide`](https://pypi.org/project/valkey-glide/)
|
||||
client (Rust core + async Python wrapper), using its native `ft` (search) command namespace.
|
||||
|
||||
### Key Features
|
||||
|
||||
- **Vector search**: ANN via HNSW or exact via FLAT, with COSINE / L2 / IP distance metrics
|
||||
- **Full-text search**: term, prefix and phrase matching over indexed text fields
|
||||
- **Hybrid search**: a metadata/text filter pre-selects candidates, then KNN ranks them
|
||||
- **In-memory speed**: vectors and documents are stored as Valkey HASH keys
|
||||
- **Auth + TLS**: optional username/password and TLS for production (toB / SaaS) deployments
|
||||
|
||||
### Licensing
|
||||
|
||||
- Valkey core and the Search module are **BSD-3-Clause**.
|
||||
- The `valkey-glide` client is **Apache-2.0**.
|
||||
|
||||
Both are compatible with LangBot.
|
||||
|
||||
## Installation
|
||||
|
||||
Valkey Search support is included when you install LangBot — the `valkey-glide` dependency is
|
||||
declared in `pyproject.toml`. To install manually:
|
||||
|
||||
```bash
|
||||
pip install 'valkey-glide>=2.4.1,<3.0.0'
|
||||
```
|
||||
|
||||
You also need a running Valkey server with the Search module loaded. The simplest way is the
|
||||
bundled image:
|
||||
|
||||
```bash
|
||||
# Run valkey-bundle (includes the Search module) on host port 6380
|
||||
podman run -d --name valkey-test-langbot -p 6380:6379 valkey/valkey-bundle:9.1.0
|
||||
# (docker run ... works identically)
|
||||
```
|
||||
|
||||
`valkey-bundle` ships multi-arch images (linux/amd64 + linux/arm64), so it runs on both CI
|
||||
(x86_64) and Apple-silicon dev machines.
|
||||
|
||||
## Configuration
|
||||
|
||||
Valkey Search is **opt-in and disabled by default** — the default `vdb.use` stays `chroma`,
|
||||
so existing single-process deployments are unaffected. To enable it, edit your `config.yaml`:
|
||||
|
||||
```yaml
|
||||
vdb:
|
||||
use: valkey_search
|
||||
valkey_search:
|
||||
host: 'localhost'
|
||||
port: 6379 # use 6380 if you started the container as shown above
|
||||
db: 0
|
||||
password: '' # optional (ACL / requirepass) — never logged
|
||||
username: '' # optional (ACL user)
|
||||
tls: false # optional (toB / SaaS)
|
||||
index_algorithm: 'HNSW' # HNSW | FLAT
|
||||
distance_metric: 'COSINE' # COSINE | L2 | IP
|
||||
request_timeout: 5000 # per-request timeout in ms
|
||||
```
|
||||
|
||||
| Option | Default | Description |
|
||||
|--------|---------|-------------|
|
||||
| `host` | `localhost` | Valkey host |
|
||||
| `port` | `6379` | Valkey port |
|
||||
| `db` | `0` | Logical database id |
|
||||
| `password` | `''` | Optional auth password (empty = no auth). Never logged. |
|
||||
| `username` | `''` | Optional ACL username. Configuring a username without a password fails closed (raises) rather than connecting unauthenticated. |
|
||||
| `tls` | `false` | Enable TLS for the connection |
|
||||
| `index_algorithm` | `HNSW` | `HNSW` (approximate) or `FLAT` (exact) |
|
||||
| `distance_metric` | `COSINE` | `COSINE`, `L2`, or `IP` |
|
||||
| `request_timeout` | `5000` | Per-request timeout in milliseconds. The valkey-glide default (250ms) is too low for vector KNN under load; raise it further for remote/cross-AZ Valkey. |
|
||||
|
||||
### Connection behavior
|
||||
|
||||
The backend uses a **lazy** connection (`lazy_connect=True`): the client is created on first
|
||||
use and the connection is deferred to the first command. A misconfigured or unreachable Valkey
|
||||
server therefore does **not** block LangBot from booting — knowledge-base operations will error
|
||||
at call time instead, and you can recover by switching `vdb.use` back to another backend.
|
||||
|
||||
The connection sets a fixed `client_name` of `langbot_vector_client` so it is identifiable in
|
||||
`CLIENT LIST` and monitoring dashboards.
|
||||
|
||||
## Supported search types
|
||||
|
||||
| Type | Behavior |
|
||||
|------|----------|
|
||||
| `vector` | Pure KNN over the embedding field |
|
||||
| `full_text` | Term/phrase match over the indexed `document` text field |
|
||||
| `hybrid` | Metadata/text filter **pre-selects** candidates, then KNN ranks them |
|
||||
|
||||
### ⚠️ Important: `vector_weight` is NOT honored
|
||||
|
||||
Valkey Search hybrid queries follow a **filter-then-KNN** model: the filter (and/or full-text
|
||||
clause) narrows the candidate set, and the KNN stage ranks the survivors by vector distance.
|
||||
There is **no native weighted score fusion** (unlike, e.g., SeekDB's RRF boost).
|
||||
|
||||
For interface compatibility the backend still accepts a `vector_weight` argument, but it is
|
||||
**ignored** — passing different weights does not change result ordering. The first time a
|
||||
non-default weight is supplied, the backend logs a one-time warning.
|
||||
|
||||
If weighted hybrid ranking is needed in the future, it can be added **application-side** (run
|
||||
vector KNN and full-text search separately and blend the scores). That is intentionally out of
|
||||
scope for this integration.
|
||||
|
||||
## Metadata & filtering
|
||||
|
||||
Documents are stored as Valkey HASH keys under the prefix `kb:{collection}:{id}` with fields:
|
||||
|
||||
- `vector` — the embedding, packed as little-endian FLOAT32
|
||||
- `document` — the raw text (indexed as TEXT for full-text/hybrid search)
|
||||
- `file_id` — promoted to an indexed TAG field so it is filterable
|
||||
- `metadata_json` — the full metadata dict, preserved verbatim as JSON
|
||||
|
||||
Only **indexed** fields are filterable. Currently that is `file_id`. Filters referencing
|
||||
non-indexed metadata keys are dropped with a warning (the same pragmatism used by the Milvus
|
||||
and pgvector backends). All other metadata still round-trips intact via `metadata_json`.
|
||||
|
||||
Supported filter operators (canonical Chroma-style `where` syntax): `$eq`, `$ne`, `$gt`,
|
||||
`$gte`, `$lt`, `$lte`, `$in`, `$nin`. Multiple top-level keys are AND-ed.
|
||||
|
||||
## Testing
|
||||
|
||||
Unit tests (filter mapping, float32 packing, reply parsing, import guard) run in the fast lane
|
||||
with no server:
|
||||
|
||||
```bash
|
||||
uv run pytest tests/unit_tests/vector/test_valkey_search_filter.py -q
|
||||
```
|
||||
|
||||
Integration tests are **slow-gated** on `TEST_VALKEY_URL` and require a running server:
|
||||
|
||||
```bash
|
||||
podman run -d --name valkey-test-langbot -p 6380:6379 valkey/valkey-bundle:9.1.0
|
||||
TEST_VALKEY_URL=valkey://localhost:6380 \
|
||||
uv run pytest tests/integration/vector/test_valkey_search.py -m slow -q
|
||||
```
|
||||
|
||||
The default upstream fast CI lane (`-m "not slow"`) skips these, matching the existing
|
||||
PostgreSQL migration-test precedent.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
| Symptom | Cause / fix |
|
||||
|---------|-------------|
|
||||
| Tests skip with "Valkey Search module not available" | The server is plain Valkey without the Search module. Use the `valkey/valkey-bundle` image. |
|
||||
| `ConnectionError` at call time | Check `host`/`port`/auth; remember `lazy_connect` defers errors to first use. |
|
||||
| Empty search results right after insert | The Search indexer is asynchronous; results become visible within a short delay. The integration tests poll/retry to account for this. |
|
||||
| Hybrid ranking ignores `vector_weight` | Expected — see the caveat above. |
|
||||
|
||||
## Production considerations
|
||||
|
||||
- **Cluster mode**: Valkey Search in cluster mode uses an additional coordination port. This
|
||||
integration targets standalone mode; cluster support is a future consideration.
|
||||
- **Persistence**: configure Valkey RDB/AOF persistence if the knowledge base must survive
|
||||
restarts; otherwise an in-memory store is ephemeral.
|
||||
- **Security**: set `password`/`username` and `tls: true` for any non-local deployment.
|
||||
Credentials are never written to logs.
|
||||
@@ -0,0 +1,713 @@
|
||||
# Workflow 系统开发者文档
|
||||
|
||||
本文档面向 LangBot 开发者,详细介绍 Workflow 系统的技术架构、核心组件和扩展方法。
|
||||
|
||||
## 目录
|
||||
|
||||
- [系统架构概述](#系统架构概述)
|
||||
- [目录结构](#目录结构)
|
||||
- [核心组件](#核心组件)
|
||||
- [后端模块](#后端模块)
|
||||
- [前端组件](#前端组件)
|
||||
- [数据库表结构](#数据库表结构)
|
||||
- [API 接口文档](#api-接口文档)
|
||||
- [如何添加新节点类型](#如何添加新节点类型)
|
||||
- [调试功能实现](#调试功能实现)
|
||||
|
||||
---
|
||||
|
||||
## 系统架构概述
|
||||
|
||||
Workflow 系统采用前后端分离架构,主要包含以下层次:
|
||||
|
||||
```
|
||||
┌─────────────────────────────────────────────────────────────┐
|
||||
│ 前端层 (React) │
|
||||
│ ┌─────────────┬──────────────┬──────────────┬───────────┐ │
|
||||
│ │ 可视化编辑器 │ 节点面板 │ 属性面板 │ 调试器 │ │
|
||||
│ │ ReactFlow │ NodePalette │ PropertyPanel│ Debugger │ │
|
||||
│ └─────────────┴──────────────┴──────────────┴───────────┘ │
|
||||
├─────────────────────────────────────────────────────────────┤
|
||||
│ API 层 (Quart) │
|
||||
│ ┌─────────────┬──────────────┬──────────────────────────┐ │
|
||||
│ │ Workflow API│ Debug API │ Node Types API │ │
|
||||
│ └─────────────┴──────────────┴──────────────────────────┘ │
|
||||
├─────────────────────────────────────────────────────────────┤
|
||||
│ 核心引擎层 (Python) │
|
||||
│ ┌─────────────┬──────────────┬──────────────┬───────────┐ │
|
||||
│ │ Executor │ Registry │ Node │ Entities │ │
|
||||
│ │ 执行引擎 │ 节点注册表 │ 节点基类 │ 数据结构 │ │
|
||||
│ └─────────────┴──────────────┴──────────────┴───────────┘ │
|
||||
├─────────────────────────────────────────────────────────────┤
|
||||
│ 存储层 (SQLAlchemy) │
|
||||
│ ┌─────────────┬──────────────┬──────────────────────────┐ │
|
||||
│ │ Workflow │ Executions │ Triggers │ │
|
||||
│ └─────────────┴──────────────┴──────────────────────────┘ │
|
||||
└─────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 目录结构
|
||||
|
||||
### 后端代码结构
|
||||
|
||||
```
|
||||
LangBot/src/langbot/pkg/
|
||||
├── workflow/ # Workflow 核心模块
|
||||
│ ├── __init__.py # 模块初始化,导出公共接口
|
||||
│ ├── entities.py # 数据实体定义
|
||||
│ ├── executor.py # 执行引擎
|
||||
│ ├── node.py # 节点基类和装饰器
|
||||
│ ├── registry.py # 节点类型注册表
|
||||
│ └── nodes/ # 内置节点实现
|
||||
│ ├── __init__.py # 注册所有内置节点
|
||||
│ ├── trigger.py # 触发节点
|
||||
│ ├── process.py # 处理节点
|
||||
│ ├── control.py # 控制节点
|
||||
│ └── action.py # 动作节点
|
||||
├── entity/persistence/
|
||||
│ └── workflow.py # 数据库模型
|
||||
├── api/http/
|
||||
│ ├── controller/groups/workflows/
|
||||
│ │ └── workflows.py # API 路由控制器
|
||||
│ └── service/
|
||||
│ └── workflow.py # 业务逻辑服务
|
||||
└── persistence/migrations/
|
||||
└── dbm026_workflow_tables.py # 数据库迁移
|
||||
```
|
||||
|
||||
### 前端代码结构
|
||||
|
||||
```
|
||||
LangBot/web/src/app/home/workflows/
|
||||
├── page.tsx # Workflow 列表页
|
||||
├── WorkflowDetailContent.tsx # 详情页内容
|
||||
├── store/
|
||||
│ └── useWorkflowStore.ts # Zustand 状态管理
|
||||
└── components/
|
||||
├── workflow-editor/ # 可视化编辑器
|
||||
│ ├── index.ts # 导出
|
||||
│ ├── WorkflowEditorComponent.tsx # 主编辑器组件
|
||||
│ ├── WorkflowNodeComponent.tsx # 自定义节点组件
|
||||
│ ├── NodePalette.tsx # 节点面板
|
||||
│ ├── PropertyPanel.tsx # 属性面板
|
||||
│ └── node-configs/ # 节点配置元数据
|
||||
│ ├── types.ts # 配置类型定义
|
||||
│ ├── trigger-configs.ts
|
||||
│ ├── ai-configs.ts
|
||||
│ ├── process-configs.ts
|
||||
│ ├── control-configs.ts
|
||||
│ ├── action-configs.ts
|
||||
│ ├── integration-configs.ts
|
||||
│ └── index.ts # 配置汇总
|
||||
├── workflow-debugger/ # 调试器组件
|
||||
│ ├── index.ts
|
||||
│ └── WorkflowDebugger.tsx
|
||||
├── workflow-form/ # 表单组件
|
||||
│ └── WorkflowFormComponent.tsx
|
||||
└── workflow-executions/ # 执行历史组件
|
||||
└── WorkflowExecutionsTab.tsx
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 核心组件
|
||||
|
||||
### 后端模块
|
||||
|
||||
#### 1. 执行引擎 (WorkflowExecutor)
|
||||
|
||||
位置:[`executor.py`](../../src/langbot/pkg/workflow/executor.py)
|
||||
|
||||
执行引擎负责工作流的实际执行,包括:
|
||||
|
||||
- **拓扑排序**:确定节点执行顺序
|
||||
- **节点执行**:调用各节点的 execute 方法
|
||||
- **控制流处理**:处理条件分支、循环、并行执行
|
||||
- **错误处理**:支持重试机制
|
||||
|
||||
```python
|
||||
class WorkflowExecutor:
|
||||
async def execute(
|
||||
self,
|
||||
workflow: WorkflowDefinition,
|
||||
context: ExecutionContext,
|
||||
start_node_id: Optional[str] = None
|
||||
) -> ExecutionContext:
|
||||
"""执行工作流"""
|
||||
# 1. 构建执行图
|
||||
# 2. 初始化节点状态
|
||||
# 3. 找到起始节点
|
||||
# 4. 按拓扑顺序执行
|
||||
```
|
||||
|
||||
**调试执行器 (DebugWorkflowExecutor)**
|
||||
|
||||
继承自 WorkflowExecutor,增加了调试支持:
|
||||
|
||||
- 断点支持
|
||||
- 单步执行
|
||||
- 暂停/继续
|
||||
- 实时日志
|
||||
|
||||
```python
|
||||
class DebugWorkflowExecutor(WorkflowExecutor):
|
||||
async def execute_debug(
|
||||
self,
|
||||
workflow: WorkflowDefinition,
|
||||
context: ExecutionContext,
|
||||
debug_state: DebugExecutionState,
|
||||
) -> ExecutionContext:
|
||||
"""调试模式执行"""
|
||||
```
|
||||
|
||||
#### 2. 节点注册表 (NodeTypeRegistry)
|
||||
|
||||
位置:[`registry.py`](../../src/langbot/pkg/workflow/registry.py)
|
||||
|
||||
单例模式管理所有节点类型:
|
||||
|
||||
```python
|
||||
class NodeTypeRegistry:
|
||||
_instance: Optional['NodeTypeRegistry'] = None
|
||||
|
||||
def register(self, node_type: str, node_class: type[WorkflowNode]):
|
||||
"""注册节点类型"""
|
||||
|
||||
def create_instance(self, node_type: str, node_id: str, config: dict) -> WorkflowNode:
|
||||
"""创建节点实例"""
|
||||
|
||||
def list_all(self) -> list[dict]:
|
||||
"""获取所有节点类型的 Schema"""
|
||||
```
|
||||
|
||||
#### 3. 节点基类 (WorkflowNode)
|
||||
|
||||
位置:[`node.py`](../../src/langbot/pkg/workflow/node.py)
|
||||
|
||||
所有节点必须继承此基类:
|
||||
|
||||
```python
|
||||
class WorkflowNode(abc.ABC):
|
||||
# 节点元数据
|
||||
type_name: str = ""
|
||||
name: str = ""
|
||||
description: str = ""
|
||||
category: str = "misc"
|
||||
icon: str = ""
|
||||
|
||||
# 端口定义
|
||||
inputs: list[NodePort] = []
|
||||
outputs: list[NodePort] = []
|
||||
|
||||
# 配置 Schema
|
||||
config_schema: list[NodeConfig] = []
|
||||
|
||||
@abc.abstractmethod
|
||||
async def execute(
|
||||
self,
|
||||
inputs: dict[str, Any],
|
||||
context: ExecutionContext
|
||||
) -> dict[str, Any]:
|
||||
"""执行节点逻辑"""
|
||||
pass
|
||||
```
|
||||
|
||||
#### 4. 数据实体 (entities.py)
|
||||
|
||||
主要数据结构:
|
||||
|
||||
```python
|
||||
class WorkflowDefinition:
|
||||
"""工作流定义"""
|
||||
uuid: str
|
||||
name: str
|
||||
nodes: list[NodeDefinition]
|
||||
edges: list[EdgeDefinition]
|
||||
settings: WorkflowSettings
|
||||
|
||||
class ExecutionContext:
|
||||
"""执行上下文"""
|
||||
execution_id: str
|
||||
workflow_id: str
|
||||
status: ExecutionStatus
|
||||
variables: dict
|
||||
node_states: dict[str, NodeState]
|
||||
history: list[ExecutionStep]
|
||||
```
|
||||
|
||||
### 前端组件
|
||||
|
||||
#### 1. WorkflowEditorComponent
|
||||
|
||||
主编辑器组件,基于 React Flow 实现:
|
||||
|
||||
- **画布交互**:拖拽、缩放、平移
|
||||
- **节点连接**:自动验证端口类型
|
||||
- **撤销/重做**:基于历史记录栈
|
||||
- **复制/粘贴**:支持多选复制
|
||||
|
||||
关键功能:
|
||||
|
||||
```tsx
|
||||
function WorkflowEditorInner() {
|
||||
const { nodes, edges, onNodesChange, onEdgesChange, onConnect } = useWorkflowStore();
|
||||
|
||||
// 拖放添加节点
|
||||
const onDrop = useCallback((event: React.DragEvent) => {
|
||||
const type = event.dataTransfer.getData('application/reactflow');
|
||||
const position = screenToFlowPosition({ x: event.clientX, y: event.clientY });
|
||||
addNode(type, position);
|
||||
}, []);
|
||||
|
||||
// 复制粘贴
|
||||
const handleCopy = useCallback(() => { ... }, []);
|
||||
const handlePaste = useCallback(() => { ... }, []);
|
||||
}
|
||||
```
|
||||
|
||||
#### 2. NodePalette
|
||||
|
||||
节点面板组件,展示可用节点类型:
|
||||
|
||||
```tsx
|
||||
function NodePalette() {
|
||||
// 按类别组织节点
|
||||
const categories = [
|
||||
{ id: 'trigger', name: '触发节点', icon: Zap },
|
||||
{ id: 'ai', name: 'AI 节点', icon: Brain },
|
||||
{ id: 'process', name: '处理节点', icon: Cpu },
|
||||
{ id: 'control', name: '控制节点', icon: GitBranch },
|
||||
{ id: 'action', name: '动作节点', icon: Send },
|
||||
{ id: 'integration', name: '集成节点', icon: Plug },
|
||||
];
|
||||
|
||||
// 拖拽开始
|
||||
const onDragStart = (event: React.DragEvent, nodeType: string) => {
|
||||
event.dataTransfer.setData('application/reactflow', nodeType);
|
||||
};
|
||||
}
|
||||
```
|
||||
|
||||
#### 3. PropertyPanel
|
||||
|
||||
属性面板组件,动态渲染节点配置表单:
|
||||
|
||||
```tsx
|
||||
function PropertyPanel() {
|
||||
const { selectedNodeId, nodes, updateNodeData } = useWorkflowStore();
|
||||
|
||||
// 根据节点类型获取配置元数据
|
||||
const selectedNode = nodes.find(n => n.id === selectedNodeId);
|
||||
const nodeConfig = getNodeConfig(selectedNode?.data?.nodeType);
|
||||
|
||||
// 动态渲染配置字段
|
||||
return (
|
||||
<div>
|
||||
{nodeConfig?.fields.map(field => (
|
||||
<ConfigField key={field.name} field={field} />
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
```
|
||||
|
||||
#### 4. WorkflowDebugger
|
||||
|
||||
调试器组件,支持实时调试:
|
||||
|
||||
```tsx
|
||||
function WorkflowDebugger({ workflowUuid, workflow }) {
|
||||
const [debugState, setDebugState] = useState<DebugState>('idle');
|
||||
const [executionId, setExecutionId] = useState<string>('');
|
||||
const [logs, setLogs] = useState<ExecutionLog[]>([]);
|
||||
|
||||
// 启动调试
|
||||
const startDebug = async () => {
|
||||
const result = await backendClient.post(
|
||||
`/api/v1/workflows/${workflowUuid}/debug/start`,
|
||||
{ context, variables, breakpoints }
|
||||
);
|
||||
setExecutionId(result.execution_id);
|
||||
};
|
||||
|
||||
// 轮询状态
|
||||
useEffect(() => {
|
||||
if (debugState === 'running') {
|
||||
const interval = setInterval(fetchState, 500);
|
||||
return () => clearInterval(interval);
|
||||
}
|
||||
}, [debugState]);
|
||||
}
|
||||
```
|
||||
|
||||
#### 5. useWorkflowStore
|
||||
|
||||
Zustand 状态管理:
|
||||
|
||||
```typescript
|
||||
interface WorkflowState {
|
||||
nodes: WorkflowNode[];
|
||||
edges: WorkflowEdge[];
|
||||
selectedNodeId: string | null;
|
||||
history: HistoryEntry[];
|
||||
historyIndex: number;
|
||||
isDirty: boolean;
|
||||
|
||||
// Actions
|
||||
addNode: (type: string, position: XYPosition) => void;
|
||||
updateNodeData: (nodeId: string, data: Partial<NodeData>) => void;
|
||||
deleteNode: (nodeId: string) => void;
|
||||
undo: () => void;
|
||||
redo: () => void;
|
||||
}
|
||||
|
||||
export const useWorkflowStore = create<WorkflowState>((set, get) => ({
|
||||
// ... state and actions
|
||||
}));
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 数据库表结构
|
||||
|
||||
### workflows 表
|
||||
|
||||
```sql
|
||||
CREATE TABLE workflows (
|
||||
uuid VARCHAR(255) PRIMARY KEY,
|
||||
name VARCHAR(255) NOT NULL,
|
||||
description TEXT,
|
||||
emoji VARCHAR(10) DEFAULT '🔄',
|
||||
version INTEGER DEFAULT 1,
|
||||
is_enabled BOOLEAN DEFAULT TRUE,
|
||||
definition JSON NOT NULL, -- 节点和边定义
|
||||
global_config JSON DEFAULT '{}', -- 全局配置
|
||||
extensions_preferences JSON, -- 插件和 MCP 配置
|
||||
created_at TIMESTAMP,
|
||||
updated_at TIMESTAMP
|
||||
);
|
||||
```
|
||||
|
||||
### workflow_versions 表
|
||||
|
||||
```sql
|
||||
CREATE TABLE workflow_versions (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
workflow_uuid VARCHAR(255) NOT NULL,
|
||||
version INTEGER NOT NULL,
|
||||
definition JSON NOT NULL,
|
||||
global_config JSON DEFAULT '{}',
|
||||
created_at TIMESTAMP,
|
||||
created_by VARCHAR(255),
|
||||
UNIQUE(workflow_uuid, version)
|
||||
);
|
||||
```
|
||||
|
||||
### workflow_executions 表
|
||||
|
||||
```sql
|
||||
CREATE TABLE workflow_executions (
|
||||
uuid VARCHAR(255) PRIMARY KEY,
|
||||
workflow_uuid VARCHAR(255) NOT NULL,
|
||||
workflow_version INTEGER NOT NULL,
|
||||
status VARCHAR(20) NOT NULL, -- pending/running/completed/failed/cancelled
|
||||
trigger_type VARCHAR(50),
|
||||
trigger_data JSON,
|
||||
variables JSON,
|
||||
start_time TIMESTAMP,
|
||||
end_time TIMESTAMP,
|
||||
error TEXT,
|
||||
created_at TIMESTAMP
|
||||
);
|
||||
```
|
||||
|
||||
### workflow_node_executions 表
|
||||
|
||||
```sql
|
||||
CREATE TABLE workflow_node_executions (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
execution_uuid VARCHAR(255) NOT NULL,
|
||||
node_id VARCHAR(100) NOT NULL,
|
||||
node_type VARCHAR(50) NOT NULL,
|
||||
status VARCHAR(20) NOT NULL,
|
||||
inputs JSON,
|
||||
outputs JSON,
|
||||
start_time TIMESTAMP,
|
||||
end_time TIMESTAMP,
|
||||
error TEXT,
|
||||
retry_count INTEGER DEFAULT 0
|
||||
);
|
||||
```
|
||||
|
||||
### workflow_triggers 表
|
||||
|
||||
```sql
|
||||
CREATE TABLE workflow_triggers (
|
||||
uuid VARCHAR(255) PRIMARY KEY,
|
||||
workflow_uuid VARCHAR(255) NOT NULL,
|
||||
type VARCHAR(50) NOT NULL, -- message/cron/event/webhook
|
||||
config JSON NOT NULL,
|
||||
is_enabled BOOLEAN DEFAULT TRUE,
|
||||
priority INTEGER DEFAULT 0,
|
||||
created_at TIMESTAMP,
|
||||
updated_at TIMESTAMP
|
||||
);
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## API 接口文档
|
||||
|
||||
### Workflow CRUD
|
||||
|
||||
| 方法 | 路径 | 描述 |
|
||||
|-----|------|------|
|
||||
| GET | `/api/v1/workflows` | 获取工作流列表 |
|
||||
| POST | `/api/v1/workflows` | 创建工作流 |
|
||||
| GET | `/api/v1/workflows/:uuid` | 获取单个工作流 |
|
||||
| PUT | `/api/v1/workflows/:uuid` | 更新工作流 |
|
||||
| DELETE | `/api/v1/workflows/:uuid` | 删除工作流 |
|
||||
| POST | `/api/v1/workflows/:uuid/copy` | 复制工作流 |
|
||||
|
||||
### 执行相关
|
||||
|
||||
| 方法 | 路径 | 描述 |
|
||||
|-----|------|------|
|
||||
| POST | `/api/v1/workflows/:uuid/execute` | 手动执行工作流 |
|
||||
| GET | `/api/v1/workflows/:uuid/executions` | 获取执行记录 |
|
||||
|
||||
### 版本管理
|
||||
|
||||
| 方法 | 路径 | 描述 |
|
||||
|-----|------|------|
|
||||
| GET | `/api/v1/workflows/:uuid/versions` | 获取版本列表 |
|
||||
| POST | `/api/v1/workflows/:uuid/rollback/:version` | 回滚到指定版本 |
|
||||
|
||||
### 调试 API
|
||||
|
||||
| 方法 | 路径 | 描述 |
|
||||
|-----|------|------|
|
||||
| POST | `/api/v1/workflows/:uuid/debug/start` | 启动调试 |
|
||||
| POST | `/api/v1/workflows/:uuid/debug/:exec_id/pause` | 暂停执行 |
|
||||
| POST | `/api/v1/workflows/:uuid/debug/:exec_id/resume` | 继续执行 |
|
||||
| POST | `/api/v1/workflows/:uuid/debug/:exec_id/stop` | 停止执行 |
|
||||
| POST | `/api/v1/workflows/:uuid/debug/:exec_id/step` | 单步执行 |
|
||||
| GET | `/api/v1/workflows/:uuid/debug/:exec_id/state` | 获取调试状态 |
|
||||
|
||||
### 节点类型
|
||||
|
||||
| 方法 | 路径 | 描述 |
|
||||
|-----|------|------|
|
||||
| GET | `/api/v1/workflows/_/node-types` | 获取所有节点类型 |
|
||||
| GET | `/api/v1/workflows/_/node-types/categories` | 按类别获取节点类型 |
|
||||
|
||||
---
|
||||
|
||||
## 如何添加新节点类型
|
||||
|
||||
### 步骤 1:创建节点类
|
||||
|
||||
在 `LangBot/src/langbot/pkg/workflow/nodes/` 下创建或修改文件:
|
||||
|
||||
```python
|
||||
from ..node import WorkflowNode, NodePort, NodeConfig, workflow_node
|
||||
from ..entities import ExecutionContext
|
||||
|
||||
@workflow_node('my_custom_node')
|
||||
class MyCustomNode(WorkflowNode):
|
||||
"""自定义节点"""
|
||||
|
||||
# 元数据
|
||||
type_name = 'my_custom_node'
|
||||
name = '我的自定义节点'
|
||||
description = '这是一个自定义节点'
|
||||
category = 'process' # trigger/process/control/action/integration
|
||||
icon = '🔧'
|
||||
|
||||
# 输入端口
|
||||
inputs = [
|
||||
NodePort(name='input', type='string', description='输入数据', required=True),
|
||||
]
|
||||
|
||||
# 输出端口
|
||||
outputs = [
|
||||
NodePort(name='output', type='string', description='输出数据'),
|
||||
]
|
||||
|
||||
# 配置字段
|
||||
config_schema = [
|
||||
NodeConfig(
|
||||
name='option',
|
||||
type='select',
|
||||
required=True,
|
||||
options=['选项A', '选项B'],
|
||||
description='选择一个选项'
|
||||
),
|
||||
NodeConfig(
|
||||
name='value',
|
||||
type='string',
|
||||
required=False,
|
||||
default='默认值',
|
||||
description='配置值'
|
||||
),
|
||||
]
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
inputs: dict[str, Any],
|
||||
context: ExecutionContext
|
||||
) -> dict[str, Any]:
|
||||
"""执行节点逻辑"""
|
||||
input_data = inputs.get('input', '')
|
||||
option = self.get_config('option')
|
||||
value = self.get_config('value', '')
|
||||
|
||||
# 处理逻辑
|
||||
result = f"处理: {input_data} with {option} and {value}"
|
||||
|
||||
return {'output': result}
|
||||
```
|
||||
|
||||
### 步骤 2:注册节点
|
||||
|
||||
在 `LangBot/src/langbot/pkg/workflow/nodes/__init__.py` 中导入:
|
||||
|
||||
```python
|
||||
from .process import (
|
||||
CodeExecutorNode,
|
||||
HttpRequestNode,
|
||||
DataTransformNode,
|
||||
MyCustomNode, # 添加新节点
|
||||
)
|
||||
```
|
||||
|
||||
### 步骤 3:添加前端配置
|
||||
|
||||
在 `LangBot/web/src/app/home/workflows/components/workflow-editor/node-configs/` 目录下添加配置:
|
||||
|
||||
```typescript
|
||||
// process-configs.ts
|
||||
export const processNodeConfigs: NodeConfigMap = {
|
||||
// ... 其他配置
|
||||
|
||||
my_custom_node: {
|
||||
type: 'my_custom_node',
|
||||
label: 'workflows.nodes.myCustomNode',
|
||||
description: 'workflows.nodes.myCustomNodeDesc',
|
||||
icon: 'Wrench',
|
||||
category: 'process',
|
||||
fields: [
|
||||
{
|
||||
name: 'option',
|
||||
type: 'select',
|
||||
label: 'workflows.fields.option',
|
||||
required: true,
|
||||
options: [
|
||||
{ value: '选项A', label: '选项 A' },
|
||||
{ value: '选项B', label: '选项 B' },
|
||||
],
|
||||
},
|
||||
{
|
||||
name: 'value',
|
||||
type: 'string',
|
||||
label: 'workflows.fields.value',
|
||||
required: false,
|
||||
defaultValue: '默认值',
|
||||
},
|
||||
],
|
||||
},
|
||||
};
|
||||
```
|
||||
|
||||
### 步骤 4:添加国际化
|
||||
|
||||
在 `LangBot/web/src/i18n/locales/` 中添加翻译:
|
||||
|
||||
```typescript
|
||||
// zh-Hans.ts
|
||||
workflows: {
|
||||
nodes: {
|
||||
myCustomNode: '我的自定义节点',
|
||||
myCustomNodeDesc: '这是一个自定义节点',
|
||||
},
|
||||
fields: {
|
||||
option: '选项',
|
||||
value: '值',
|
||||
},
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 调试功能实现
|
||||
|
||||
### 后端调试状态管理
|
||||
|
||||
```python
|
||||
class DebugExecutionState:
|
||||
"""调试执行状态"""
|
||||
|
||||
def __init__(self, execution_id: str, breakpoints: list[str] = None):
|
||||
self.execution_id = execution_id
|
||||
self.status: str = 'running'
|
||||
self.is_paused: bool = False
|
||||
self.is_stopped: bool = False
|
||||
self.breakpoints: set[str] = set(breakpoints or [])
|
||||
self.logs: list[ExecutionLog] = []
|
||||
self._pause_event = asyncio.Event()
|
||||
|
||||
def pause(self):
|
||||
"""暂停执行"""
|
||||
self.is_paused = True
|
||||
self._pause_event.clear()
|
||||
|
||||
def resume(self):
|
||||
"""继续执行"""
|
||||
self.is_paused = False
|
||||
self._pause_event.set()
|
||||
|
||||
async def wait_if_paused(self):
|
||||
"""如果暂停则等待"""
|
||||
if self.is_paused:
|
||||
await self._pause_event.wait()
|
||||
```
|
||||
|
||||
### 前端调试流程
|
||||
|
||||
1. **设置断点**:点击节点设置断点
|
||||
2. **启动调试**:调用 `/debug/start` 启动调试执行
|
||||
3. **轮询状态**:定期调用 `/debug/:id/state` 获取状态
|
||||
4. **控制执行**:调用 pause/resume/step/stop 控制执行
|
||||
5. **查看日志**:实时显示执行日志和节点状态
|
||||
|
||||
```typescript
|
||||
// 调试状态轮询
|
||||
const fetchDebugState = async () => {
|
||||
const state = await backendClient.get(
|
||||
`/api/v1/workflows/${workflowUuid}/debug/${executionId}/state`
|
||||
);
|
||||
|
||||
// 更新节点状态
|
||||
setNodeStates(state.node_states);
|
||||
|
||||
// 追加新日志
|
||||
if (state.new_logs.length > 0) {
|
||||
setLogs(prev => [...prev, ...state.new_logs]);
|
||||
}
|
||||
|
||||
// 检查完成状态
|
||||
if (state.status === 'completed' || state.status === 'error') {
|
||||
setDebugState('idle');
|
||||
}
|
||||
};
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 扩展阅读
|
||||
|
||||
- [Workflow 功能设计文档](../../../plans/langbot-workflow-design.md)
|
||||
- [用户使用指南](../user-guide/workflow-guide.md)
|
||||
- [API 认证文档](../API_KEY_AUTH.md)
|
||||
@@ -0,0 +1,425 @@
|
||||
# Workflow 用户指南
|
||||
|
||||
本文档帮助您了解和使用 LangBot 的 Workflow(工作流)功能,通过可视化方式构建自动化的对话处理流程。
|
||||
|
||||
## 目录
|
||||
|
||||
- [功能介绍](#功能介绍)
|
||||
- [快速入门](#快速入门)
|
||||
- [节点类型说明](#节点类型说明)
|
||||
- [编辑器使用指南](#编辑器使用指南)
|
||||
- [调试功能](#调试功能)
|
||||
- [常见问题解答](#常见问题解答)
|
||||
|
||||
---
|
||||
|
||||
## 功能介绍
|
||||
|
||||
### 什么是 Workflow?
|
||||
|
||||
Workflow(工作流)是 LangBot 提供的可视化自动化编排系统。通过拖拽节点、连接边的方式,您可以:
|
||||
|
||||
- 📝 **构建复杂的对话流程**:使用条件分支、循环等控制节点
|
||||
- 🤖 **调用 AI 能力**:集成 LLM、知识库检索、参数提取
|
||||
- 🔗 **连接外部服务**:集成 Dify、n8n、Coze 等平台
|
||||
- ⚡ **自动化任务执行**:消息触发、定时触发、Webhook 触发
|
||||
|
||||
### Workflow vs Pipeline
|
||||
|
||||
| 对比项 | Pipeline | Workflow |
|
||||
|-------|----------|----------|
|
||||
| 配置方式 | 表单配置 | 可视化拖拽 |
|
||||
| 流程控制 | 线性执行 | 支持分支、循环、并行 |
|
||||
| 适用场景 | 简单对话 | 复杂流程 |
|
||||
| 学习曲线 | 低 | 中等 |
|
||||
|
||||
---
|
||||
|
||||
## 快速入门
|
||||
|
||||
### 第一步:创建 Workflow
|
||||
|
||||
1. 在侧边栏点击 **Workflow** 进入工作流列表
|
||||
2. 点击右上角 **创建工作流** 按钮
|
||||
3. 填写基本信息:
|
||||
- **名称**:给工作流起一个描述性的名字
|
||||
- **描述**:可选,说明工作流的用途
|
||||
- **图标**:选择一个 emoji 作为标识
|
||||
|
||||
### 第二步:添加节点
|
||||
|
||||
进入编辑器后,左侧是节点面板,中间是画布区域,右侧是属性面板。
|
||||
|
||||
1. **添加触发节点**:从左侧面板拖拽一个"消息触发"节点到画布
|
||||
2. **添加 AI 节点**:拖拽一个"LLM 调用"节点
|
||||
3. **添加回复节点**:拖拽一个"回复消息"节点
|
||||
|
||||
### 第三步:连接节点
|
||||
|
||||
1. 将鼠标悬停在触发节点的输出端口(右侧小圆点)
|
||||
2. 按住鼠标拖拽到 LLM 节点的输入端口(左侧小圆点)
|
||||
3. 同样方式连接 LLM 节点和回复节点
|
||||
|
||||
```
|
||||
[消息触发] ──▶ [LLM 调用] ──▶ [回复消息]
|
||||
```
|
||||
|
||||
### 第四步:配置节点
|
||||
|
||||
点击 LLM 调用节点,在右侧属性面板配置:
|
||||
|
||||
- **运行方式**:选择"本地 Agent"
|
||||
- **系统提示词**:描述 AI 的角色和行为
|
||||
- **模型**:选择要使用的 LLM 模型
|
||||
|
||||
点击回复消息节点配置:
|
||||
|
||||
- **消息内容**:设置为 `{{nodes.llm_call.outputs.response}}`(引用 LLM 输出)
|
||||
|
||||
### 第五步:保存并绑定
|
||||
|
||||
1. 点击工具栏的 **保存** 按钮
|
||||
2. 返回 Bot 配置页面
|
||||
3. 在 Bot 的绑定设置中选择 **Workflow**,然后选择刚创建的工作流
|
||||
|
||||
恭喜!您已经创建了第一个 Workflow。
|
||||
|
||||
---
|
||||
|
||||
## 节点类型说明
|
||||
|
||||
### 触发节点 (Trigger)
|
||||
|
||||
触发节点是工作流的入口,定义何时启动执行。
|
||||
|
||||
| 节点 | 说明 | 输出 |
|
||||
|-----|------|------|
|
||||
| 消息触发 | 收到消息时触发 | message, sender_id, platform |
|
||||
| 定时触发 | 按 Cron 表达式定时触发 | timestamp |
|
||||
| Webhook 触发 | 收到 HTTP 请求时触发 | request_body, headers |
|
||||
| 事件触发 | 系统事件触发 | event_type, event_data |
|
||||
|
||||
**消息触发配置示例**:
|
||||
|
||||
```yaml
|
||||
触发条件:
|
||||
- 关键词匹配: ["帮助", "help"]
|
||||
- 平台: ["wechat", "qq"]
|
||||
```
|
||||
|
||||
### AI 节点
|
||||
|
||||
AI 节点用于调用各种 AI 能力。
|
||||
|
||||
| 节点 | 说明 | 典型用途 |
|
||||
|-----|------|---------|
|
||||
| LLM 调用 | 调用大语言模型 | 生成回复、理解意图 |
|
||||
| 问题分类器 | 对用户问题分类 | 路由到不同处理分支 |
|
||||
| 参数提取器 | 从文本提取结构化数据 | 提取订单号、日期等 |
|
||||
| 知识库检索 | 查询知识库 | RAG 增强回复 |
|
||||
|
||||
**LLM 调用配置示例**:
|
||||
|
||||
```yaml
|
||||
运行方式: 本地 Agent
|
||||
模型: gpt-4
|
||||
系统提示词: |
|
||||
你是一个友好的客服助手。
|
||||
请根据用户的问题提供帮助。
|
||||
温度: 0.7
|
||||
最大 Token 数: 2000
|
||||
```
|
||||
|
||||
### 处理节点 (Process)
|
||||
|
||||
处理节点用于数据处理和外部调用。
|
||||
|
||||
| 节点 | 说明 | 典型用途 |
|
||||
|-----|------|---------|
|
||||
| 代码执行 | 执行 Python/JavaScript 代码 | 数据处理、格式转换 |
|
||||
| HTTP 请求 | 发送 HTTP 请求 | 调用外部 API |
|
||||
| 数据转换 | JSON/模板转换 | 数据格式化 |
|
||||
|
||||
**HTTP 请求配置示例**:
|
||||
|
||||
```yaml
|
||||
URL: https://api.example.com/data
|
||||
方法: POST
|
||||
请求头:
|
||||
Content-Type: application/json
|
||||
Authorization: Bearer {{variables.api_key}}
|
||||
请求体: |
|
||||
{"query": "{{message.content}}"}
|
||||
```
|
||||
|
||||
### 控制节点 (Control)
|
||||
|
||||
控制节点用于流程控制。
|
||||
|
||||
| 节点 | 说明 | 用途 |
|
||||
|-----|------|------|
|
||||
| 条件分支 | 二选一分支 | if-else 逻辑 |
|
||||
| 多路分支 | 多选一分支 | switch-case 逻辑 |
|
||||
| 循环 | 遍历数组 | 批量处理 |
|
||||
| 并行 | 同时执行多分支 | 并发处理 |
|
||||
| 等待 | 暂停执行 | 延时处理 |
|
||||
| 合并 | 合并多个分支 | 汇总结果 |
|
||||
|
||||
**条件分支配置示例**:
|
||||
|
||||
```yaml
|
||||
条件表达式: "{{nodes.classifier.outputs.category}}" == "complaint"
|
||||
真分支: 投诉处理
|
||||
假分支: 普通咨询
|
||||
```
|
||||
|
||||
### 动作节点 (Action)
|
||||
|
||||
动作节点执行具体操作。
|
||||
|
||||
| 节点 | 说明 | 用途 |
|
||||
|-----|------|------|
|
||||
| 发送消息 | 主动发送消息 | 通知、推送 |
|
||||
| 回复消息 | 回复当前消息 | 对话回复 |
|
||||
| 存储数据 | 保存数据到存储 | 持久化 |
|
||||
| 调用 Pipeline | 调用现有 Pipeline | 复用现有流程 |
|
||||
|
||||
**回复消息配置示例**:
|
||||
|
||||
```yaml
|
||||
消息内容: |
|
||||
感谢您的咨询!
|
||||
|
||||
{{nodes.llm_call.outputs.response}}
|
||||
|
||||
如有其他问题,随时联系我。
|
||||
```
|
||||
|
||||
### 集成节点 (Integration)
|
||||
|
||||
集成节点连接外部平台。
|
||||
|
||||
| 节点 | 说明 | 平台 |
|
||||
|-----|------|------|
|
||||
| Dify 工作流 | 调用 Dify 应用 | Dify |
|
||||
| Dify 知识库 | 查询 Dify 知识库 | Dify |
|
||||
| n8n 工作流 | 调用 n8n 流程 | n8n |
|
||||
| Langflow | 调用 Langflow 流程 | Langflow |
|
||||
| Coze Bot | 调用扣子 Bot | Coze |
|
||||
|
||||
**Dify 工作流配置示例**:
|
||||
|
||||
```yaml
|
||||
API 地址: https://api.dify.ai/v1
|
||||
API Key: sk-xxxxx
|
||||
应用类型: workflow
|
||||
同步对话历史: true
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 编辑器使用指南
|
||||
|
||||
### 画布操作
|
||||
|
||||
| 操作 | 方式 |
|
||||
|-----|------|
|
||||
| 平移画布 | 按住鼠标中键/空格+左键 拖拽 |
|
||||
| 缩放画布 | 鼠标滚轮 / 工具栏按钮 |
|
||||
| 框选多个节点 | 按住 Shift + 拖拽框选 |
|
||||
| 适应视图 | 点击工具栏"适应"按钮 |
|
||||
|
||||
### 节点操作
|
||||
|
||||
| 操作 | 方式 |
|
||||
|-----|------|
|
||||
| 添加节点 | 从左侧面板拖拽到画布 |
|
||||
| 移动节点 | 点击节点拖拽 |
|
||||
| 删除节点 | 选中后按 Delete / 点击工具栏删除 |
|
||||
| 复制节点 | 选中后 Ctrl+C / 工具栏复制 |
|
||||
| 粘贴节点 | Ctrl+V / 工具栏粘贴 |
|
||||
|
||||
### 连接操作
|
||||
|
||||
| 操作 | 方式 |
|
||||
|-----|------|
|
||||
| 创建连接 | 从输出端口拖拽到输入端口 |
|
||||
| 删除连接 | 点击连接线后按 Delete |
|
||||
| 选中连接 | 点击连接线 |
|
||||
|
||||
### 快捷键
|
||||
|
||||
| 快捷键 | 功能 |
|
||||
|-------|------|
|
||||
| Ctrl + Z | 撤销 |
|
||||
| Ctrl + Shift + Z | 重做 |
|
||||
| Ctrl + C | 复制 |
|
||||
| Ctrl + V | 粘贴 |
|
||||
| Delete | 删除选中 |
|
||||
| Ctrl + S | 保存 |
|
||||
|
||||
### 工具栏功能
|
||||
|
||||
```
|
||||
[撤销] [重做] | [放大] [缩小] [适应] | [复制] [粘贴] [删除] | [保存] [调试]
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 调试功能
|
||||
|
||||
### 启动调试
|
||||
|
||||
1. 点击工具栏的 **调试** 按钮
|
||||
2. 在调试面板中配置初始数据:
|
||||
- **输入消息**:模拟用户发送的消息
|
||||
- **会话 ID**:可选,用于测试会话变量
|
||||
- **变量**:设置初始变量值
|
||||
|
||||
3. 点击 **开始调试** 按钮
|
||||
|
||||
### 调试控制
|
||||
|
||||
| 按钮 | 功能 |
|
||||
|-----|------|
|
||||
| ▶️ 开始/继续 | 开始或继续执行 |
|
||||
| ⏸️ 暂停 | 暂停执行 |
|
||||
| ⏹️ 停止 | 停止执行 |
|
||||
| ⏭️ 单步 | 执行下一个节点 |
|
||||
|
||||
### 断点
|
||||
|
||||
- **设置断点**:点击节点上的断点图标
|
||||
- **断点触发**:执行到断点时自动暂停
|
||||
- **查看状态**:在暂停时查看节点的输入输出
|
||||
|
||||
### 执行日志
|
||||
|
||||
调试面板下方显示实时日志:
|
||||
|
||||
```
|
||||
[INFO] 2024-01-15 10:30:00 - Starting debug execution
|
||||
[INFO] 2024-01-15 10:30:00 - Executing node: message_trigger
|
||||
[DEBUG] 2024-01-15 10:30:00 - Node inputs: {"message": "你好"}
|
||||
[INFO] 2024-01-15 10:30:01 - Node completed in 50ms
|
||||
[INFO] 2024-01-15 10:30:01 - Executing node: llm_call
|
||||
...
|
||||
```
|
||||
|
||||
### 节点状态颜色
|
||||
|
||||
| 颜色 | 状态 |
|
||||
|-----|------|
|
||||
| 灰色 | 待执行 |
|
||||
| 蓝色 | 执行中 |
|
||||
| 绿色 | 已完成 |
|
||||
| 红色 | 失败 |
|
||||
| 黄色 | 已跳过 |
|
||||
|
||||
---
|
||||
|
||||
## 常见问题解答
|
||||
|
||||
### Q1:如何在节点间传递数据?
|
||||
|
||||
使用表达式语法引用其他节点的输出:
|
||||
|
||||
```
|
||||
{{nodes.节点ID.outputs.输出名称}}
|
||||
```
|
||||
|
||||
例如:
|
||||
- `{{nodes.llm_call.outputs.response}}` - 引用 LLM 节点的响应
|
||||
- `{{nodes.http_request.outputs.body}}` - 引用 HTTP 请求的响应体
|
||||
|
||||
### Q2:如何使用变量?
|
||||
|
||||
Workflow 支持三种变量类型:
|
||||
|
||||
1. **工作流变量**:`{{variables.变量名}}`
|
||||
2. **会话变量**:`{{conversation_variables.变量名}}`
|
||||
3. **消息上下文**:`{{message.content}}`、`{{message.sender_id}}`
|
||||
|
||||
### Q3:条件分支如何写条件表达式?
|
||||
|
||||
支持以下运算符:
|
||||
|
||||
- 比较:`==`, `!=`, `>`, `<`, `>=`, `<=`
|
||||
- 逻辑:`and`, `or`, `not`
|
||||
- 包含:`in`
|
||||
|
||||
示例:
|
||||
```python
|
||||
# 字符串比较
|
||||
"{{nodes.classifier.outputs.intent}}" == "purchase"
|
||||
|
||||
# 数值比较
|
||||
{{nodes.extractor.outputs.amount}} > 1000
|
||||
|
||||
# 包含检查
|
||||
"退款" in "{{message.content}}"
|
||||
```
|
||||
|
||||
### Q4:如何处理错误?
|
||||
|
||||
1. **节点级重试**:在节点配置中设置重试次数
|
||||
2. **全局错误处理**:在 Workflow 设置中配置错误处理策略
|
||||
3. **条件分支**:使用条件节点检查上一节点的状态
|
||||
|
||||
### Q5:如何查看执行历史?
|
||||
|
||||
1. 进入 Workflow 详情页
|
||||
2. 点击 **执行历史** 标签
|
||||
3. 查看每次执行的状态、耗时、输入输出
|
||||
|
||||
### Q6:Workflow 可以被多个 Bot 使用吗?
|
||||
|
||||
是的。一个 Workflow 可以被多个 Bot 绑定使用,但每个 Bot 只能绑定一个处理单元(Pipeline 或 Workflow)。
|
||||
|
||||
### Q7:如何复制现有的 Workflow?
|
||||
|
||||
在 Workflow 列表页,点击工作流卡片右上角的菜单,选择"复制"即可创建副本。
|
||||
|
||||
### Q8:支持版本回滚吗?
|
||||
|
||||
支持。每次保存都会创建新版本。在 Workflow 详情页可以查看版本历史并回滚到指定版本。
|
||||
|
||||
---
|
||||
|
||||
## 最佳实践
|
||||
|
||||
### 1. 合理命名
|
||||
|
||||
- 为节点和 Workflow 使用描述性名称
|
||||
- 使用统一的命名规范
|
||||
|
||||
### 2. 模块化设计
|
||||
|
||||
- 将复杂流程拆分为多个小 Workflow
|
||||
- 使用"调用 Pipeline"节点复用现有流程
|
||||
|
||||
### 3. 错误处理
|
||||
|
||||
- 为关键节点设置重试机制
|
||||
- 使用条件分支处理异常情况
|
||||
- 添加日志记录便于排查问题
|
||||
|
||||
### 4. 测试先行
|
||||
|
||||
- 使用调试功能充分测试
|
||||
- 准备多种测试场景
|
||||
- 检查边界情况
|
||||
|
||||
### 5. 性能优化
|
||||
|
||||
- 避免不必要的节点
|
||||
- 使用并行节点提高效率
|
||||
- 合理设置超时时间
|
||||
|
||||
---
|
||||
|
||||
## 更多资源
|
||||
|
||||
- [开发者文档](../development/workflow-system.md)
|
||||
- [设计文档](../../../plans/langbot-workflow-design.md)
|
||||
- [API 文档](../service-api-openapi.json)
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+2
-3
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "langbot"
|
||||
version = "4.10.5"
|
||||
version = "4.10.4"
|
||||
description = "Production-grade platform for building agentic IM bots"
|
||||
readme = "README.md"
|
||||
license-files = ["LICENSE"]
|
||||
@@ -70,7 +70,7 @@ dependencies = [
|
||||
"chromadb>=1.0.0,<2.0.0",
|
||||
"qdrant-client (>=1.15.1,<2.0.0)",
|
||||
"pyseekdb==1.1.0.post3",
|
||||
"langbot-plugin==0.4.13",
|
||||
"langbot-plugin @ file:///home/qinjunyan/code/projects/langbot/langbot-plugin-sdk",
|
||||
"asyncpg>=0.30.0",
|
||||
"line-bot-sdk>=3.19.0",
|
||||
"matrix-nio>=0.25.2",
|
||||
@@ -80,7 +80,6 @@ dependencies = [
|
||||
"pgvector>=0.4.1",
|
||||
"botocore>=1.42.39",
|
||||
"litellm>=1.0.0",
|
||||
"valkey-glide>=2.4.1,<3.0.0",
|
||||
]
|
||||
keywords = [
|
||||
"bot",
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -471,7 +471,7 @@ async def on_msg(event_context: context.EventContext):
|
||||
if isinstance(component, platform_message.Plain):
|
||||
text_parts.append(component.text)
|
||||
text = "".join(text_parts).strip()
|
||||
|
||||
|
||||
if should_handle(text):
|
||||
event_context.prevent_default()
|
||||
event_context.prevent_postorder()
|
||||
|
||||
@@ -109,62 +109,6 @@ class AsyncDifyServiceClient:
|
||||
if chunk.startswith('data:'):
|
||||
yield json.loads(chunk[5:])
|
||||
|
||||
async def workflow_submit(
|
||||
self,
|
||||
form_token: str,
|
||||
workflow_run_id: str,
|
||||
inputs: dict[str, typing.Any],
|
||||
user: str,
|
||||
action: str = '',
|
||||
timeout: float = 120.0,
|
||||
) -> typing.AsyncGenerator[dict[str, typing.Any], None]:
|
||||
"""Submit human input to resume a paused workflow, then stream events.
|
||||
|
||||
1. POST /form/human_input/{form_token} to submit the form
|
||||
2. GET /workflow/{task_id}/events to stream the resumed workflow events
|
||||
"""
|
||||
|
||||
headers = {
|
||||
'Authorization': f'Bearer {self.api_key}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(
|
||||
base_url=self.base_url,
|
||||
trust_env=True,
|
||||
timeout=timeout,
|
||||
) as client:
|
||||
# Step 1: Submit the form
|
||||
payload: dict[str, typing.Any] = {
|
||||
'inputs': inputs if isinstance(inputs, dict) else {},
|
||||
'user': user,
|
||||
'action': action,
|
||||
}
|
||||
|
||||
submit_resp = await client.post(
|
||||
f'/form/human_input/{form_token}',
|
||||
headers=headers,
|
||||
json=payload,
|
||||
)
|
||||
if submit_resp.status_code != 200:
|
||||
raise DifyAPIError(f'{submit_resp.status_code} {submit_resp.text}')
|
||||
|
||||
# Step 2: Stream resumed workflow events
|
||||
async with client.stream(
|
||||
'GET',
|
||||
f'/workflow/{workflow_run_id}/events',
|
||||
headers={'Authorization': f'Bearer {self.api_key}'},
|
||||
params={'user': user},
|
||||
) as r:
|
||||
if r.status_code != 200:
|
||||
body = (await r.aread()).decode(errors='replace')
|
||||
raise DifyAPIError(f'{r.status_code} {body}')
|
||||
async for chunk in r.aiter_lines():
|
||||
if chunk.strip() == '':
|
||||
continue
|
||||
if chunk.startswith('data:'):
|
||||
yield json.loads(chunk[5:])
|
||||
|
||||
async def upload_file(
|
||||
self,
|
||||
file: httpx._types.FileTypes,
|
||||
|
||||
@@ -1,48 +1,17 @@
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import typing
|
||||
import uuid
|
||||
import urllib.parse
|
||||
from typing import Awaitable, Callable, Optional
|
||||
from typing import Callable
|
||||
import dingtalk_stream # type: ignore
|
||||
import websockets
|
||||
from .EchoHandler import EchoTextHandler
|
||||
from .card_callback import DingTalkCardActionHandler
|
||||
from .dingtalkevent import DingTalkEvent
|
||||
import httpx
|
||||
import traceback
|
||||
|
||||
|
||||
_stdout_logger = logging.getLogger('langbot.dingtalk_api')
|
||||
|
||||
|
||||
DINGTALK_OPENAPI_BASE = 'https://api.dingtalk.com'
|
||||
|
||||
|
||||
def _stringify_card_param_map(card_param_map: Optional[dict]) -> dict:
|
||||
"""DingTalk cardParamMap only accepts string values.
|
||||
|
||||
Keep callers free to pass structured values for template variables such
|
||||
as button groups or select options, then encode them once at the API
|
||||
boundary.
|
||||
"""
|
||||
if not card_param_map:
|
||||
return {}
|
||||
result = {}
|
||||
for key, value in card_param_map.items():
|
||||
if value is None:
|
||||
result[key] = ''
|
||||
elif isinstance(value, str):
|
||||
result[key] = value
|
||||
else:
|
||||
result[key] = json.dumps(value, ensure_ascii=False)
|
||||
return result
|
||||
|
||||
|
||||
class DingTalkClient:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -52,7 +21,6 @@ class DingTalkClient:
|
||||
robot_code: str,
|
||||
markdown_card: bool,
|
||||
logger: None,
|
||||
card_action_callback: Optional[Callable[[dict], Awaitable[None]]] = None,
|
||||
):
|
||||
"""初始化 WebSocket 连接并自动启动"""
|
||||
self.credential = dingtalk_stream.Credential(client_id, client_secret)
|
||||
@@ -62,14 +30,6 @@ class DingTalkClient:
|
||||
# 在 DingTalkClient 中传入自己作为参数,避免循环导入
|
||||
self.EchoTextHandler = EchoTextHandler(self)
|
||||
self.client.register_callback_handler(dingtalk_stream.chatbot.ChatbotMessage.TOPIC, self.EchoTextHandler)
|
||||
# STREAM-mode card action button click handler. Forwards parsed payload
|
||||
# to the adapter so it can resume paused Dify workflows.
|
||||
self.card_action_callback = card_action_callback
|
||||
self.card_action_handler = DingTalkCardActionHandler(self.client, self._on_card_action)
|
||||
self.client.register_callback_handler(
|
||||
dingtalk_stream.handlers.CallbackHandler.TOPIC_CARD_CALLBACK,
|
||||
self.card_action_handler,
|
||||
)
|
||||
self._message_handlers = {
|
||||
'example': [],
|
||||
}
|
||||
@@ -79,24 +39,8 @@ class DingTalkClient:
|
||||
self.access_token_expiry_time = ''
|
||||
self.markdown_card = markdown_card
|
||||
self.logger = logger
|
||||
# Legacy access_token used by the OLD oapi.dingtalk.com endpoints
|
||||
# (e.g. /media/upload, which is the only documented way to get an
|
||||
# `@xxx` media_id usable in card Avatar.imageUrl). The new v1.0
|
||||
# token doesn't work there — different auth domain.
|
||||
self.legacy_access_token = ''
|
||||
self.legacy_access_token_expiry_time: typing.Optional[float] = None
|
||||
self._stopped = False # Flag to control the event loop
|
||||
|
||||
async def _on_card_action(self, payload: dict) -> None:
|
||||
"""Dispatch a parsed card-action payload to the adapter callback."""
|
||||
if self.card_action_callback is None:
|
||||
return
|
||||
try:
|
||||
await self.card_action_callback(payload)
|
||||
except Exception:
|
||||
if self.logger:
|
||||
await self.logger.error(f'DingTalk card action callback error: {traceback.format_exc()}')
|
||||
|
||||
async def get_access_token(self):
|
||||
url = 'https://api.dingtalk.com/v1.0/oauth2/accessToken'
|
||||
headers = {'Content-Type': 'application/json'}
|
||||
@@ -485,35 +429,18 @@ class DingTalkClient:
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
|
||||
# For enterprise-internal robots, robotCode == AppKey (client_id).
|
||||
# The dedicated robot_code field is only required for scenario-group
|
||||
# robots or third-party robots; fall back to client_id when empty so
|
||||
# the common single-bot setup keeps working without manual config.
|
||||
robot_code = self.robot_code or self.key
|
||||
data = {
|
||||
'robotCode': robot_code,
|
||||
'robotCode': self.robot_code,
|
||||
'userIds': [target_id],
|
||||
'msgKey': 'sampleText',
|
||||
'msgParam': json.dumps({'content': content}),
|
||||
}
|
||||
_stdout_logger.info(
|
||||
'DingTalk send_proactive_message_to_one request: robotCode=%s target_id=%s content_len=%d',
|
||||
robot_code,
|
||||
target_id,
|
||||
len(content),
|
||||
)
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(url, headers=headers, json=data)
|
||||
_stdout_logger.info(
|
||||
'DingTalk send_proactive_message_to_one response: status=%d body=%s',
|
||||
response.status_code,
|
||||
response.text[:500],
|
||||
)
|
||||
if response.status_code == 200:
|
||||
return
|
||||
except Exception:
|
||||
_stdout_logger.exception('DingTalk send_proactive_message_to_one error')
|
||||
await self.logger.error(f'failed to send proactive massage to person: {traceback.format_exc()}')
|
||||
raise Exception(f'failed to send proactive massage to person: {traceback.format_exc()}')
|
||||
|
||||
@@ -529,7 +456,7 @@ class DingTalkClient:
|
||||
}
|
||||
|
||||
data = {
|
||||
'robotCode': self.robot_code or self.key,
|
||||
'robotCode': self.robot_code,
|
||||
'openConversationId': target_id,
|
||||
'msgKey': 'sampleText',
|
||||
'msgParam': json.dumps({'content': content}),
|
||||
@@ -550,334 +477,47 @@ class DingTalkClient:
|
||||
quote_origin: bool = False,
|
||||
card_auto_layout: bool = False,
|
||||
):
|
||||
"""Create + deliver the streaming chat card for a chatbot reply.
|
||||
card_data = {}
|
||||
card_data['config'] = json.dumps({'autoLayout': card_auto_layout})
|
||||
card_data['content'] = ''
|
||||
|
||||
Replaces the old `dingtalk_stream.AICardReplier`-based path. Returns
|
||||
`(None, out_track_id)` to keep call sites compatible with the
|
||||
previous `(card_instance, card_instance_id)` shape — the first slot
|
||||
is unused now that everything is driven by out_track_id.
|
||||
"""
|
||||
out_track_id = uuid.uuid4().hex
|
||||
is_group = str(incoming_message.conversation_type) == '2'
|
||||
if is_group:
|
||||
open_space_id = f'dtv1.card//IM_GROUP.{incoming_message.conversation_id}'
|
||||
else:
|
||||
open_space_id = f'dtv1.card//IM_ROBOT.{incoming_message.sender_staff_id}'
|
||||
|
||||
card_param_map = {'content': ''}
|
||||
# 将用户的消息内容作为卡片的查询参数,方便后续处理
|
||||
if incoming_message.message_type == 'text':
|
||||
card_param_map['query'] = incoming_message.get_text_list()[0]
|
||||
card_data['query'] = incoming_message.get_text_list()[0]
|
||||
else:
|
||||
card_param_map['query'] = '...'
|
||||
card_data['query'] = '...'
|
||||
|
||||
await self.create_and_deliver_card(
|
||||
card_template_id=temp_card_id,
|
||||
out_track_id=out_track_id,
|
||||
open_space_id=open_space_id,
|
||||
is_group=is_group,
|
||||
card_param_map=card_param_map,
|
||||
card_data_config={'autoLayout': card_auto_layout},
|
||||
card_instance = dingtalk_stream.AICardReplier(self.client, incoming_message)
|
||||
# print(card_instance)
|
||||
# 先投放卡片: https://open.dingtalk.com/document/orgapp/create-and-deliver-cards
|
||||
card_instance_id = await card_instance.async_create_and_deliver_card(
|
||||
temp_card_id,
|
||||
card_data,
|
||||
)
|
||||
return None, out_track_id
|
||||
return card_instance, card_instance_id
|
||||
|
||||
async def send_card_message(self, card_instance, card_instance_id: str, content: str, is_final: bool):
|
||||
"""Stream a single chunk into an existing card's `content` field."""
|
||||
content_key = 'content'
|
||||
try:
|
||||
await self.streaming_update_card(
|
||||
out_track_id=card_instance_id,
|
||||
content_key='content',
|
||||
await card_instance.async_streaming(
|
||||
card_instance_id,
|
||||
content_key=content_key,
|
||||
content_value=content,
|
||||
append=False,
|
||||
finished=is_final,
|
||||
failed=False,
|
||||
)
|
||||
except Exception as e:
|
||||
if self.logger:
|
||||
self.logger.exception(e)
|
||||
await self.streaming_update_card(
|
||||
out_track_id=card_instance_id,
|
||||
content_key='content',
|
||||
self.logger.exception(e)
|
||||
await card_instance.async_streaming(
|
||||
card_instance_id,
|
||||
content_key=content_key,
|
||||
content_value='',
|
||||
append=False,
|
||||
finished=is_final,
|
||||
failed=True,
|
||||
)
|
||||
|
||||
async def create_and_deliver_card(
|
||||
self,
|
||||
*,
|
||||
card_template_id: str,
|
||||
out_track_id: str,
|
||||
open_space_id: str,
|
||||
is_group: bool,
|
||||
card_param_map: Optional[dict] = None,
|
||||
callback_type: str = 'STREAM',
|
||||
callback_route_key: Optional[str] = None,
|
||||
support_forward: bool = True,
|
||||
dynamic_data_source_configs: Optional[list] = None,
|
||||
card_data_config: Optional[dict] = None,
|
||||
at_user_ids: Optional[dict] = None,
|
||||
recipients: Optional[list] = None,
|
||||
) -> bool:
|
||||
"""POST /v1.0/card/instances/createAndDeliver.
|
||||
|
||||
Mirrors the SDK's `async_create_and_deliver_card` shape but exposes
|
||||
the dynamic-data-source config slot so we can register a pull URL
|
||||
for variable-length button lists.
|
||||
"""
|
||||
if not await self.check_access_token():
|
||||
await self.get_access_token()
|
||||
|
||||
cardData: dict = {'cardParamMap': _stringify_card_param_map(card_param_map)}
|
||||
if card_data_config is not None:
|
||||
cardData['config'] = json.dumps(card_data_config)
|
||||
|
||||
body: dict = {
|
||||
'cardTemplateId': card_template_id,
|
||||
'outTrackId': out_track_id,
|
||||
'cardData': cardData,
|
||||
'callbackType': callback_type,
|
||||
'openSpaceId': open_space_id,
|
||||
'imGroupOpenSpaceModel': {'supportForward': support_forward},
|
||||
'imRobotOpenSpaceModel': {'supportForward': support_forward},
|
||||
}
|
||||
if callback_type == 'HTTP' and callback_route_key:
|
||||
body['callbackRouteKey'] = callback_route_key
|
||||
|
||||
if is_group:
|
||||
deliver: dict = {'robotCode': self.robot_code or self.key}
|
||||
if at_user_ids:
|
||||
deliver['atUserIds'] = at_user_ids
|
||||
if recipients is not None:
|
||||
deliver['recipients'] = recipients
|
||||
body['imGroupOpenDeliverModel'] = deliver
|
||||
else:
|
||||
body['imRobotOpenDeliverModel'] = {'spaceType': 'IM_ROBOT'}
|
||||
|
||||
if dynamic_data_source_configs:
|
||||
body['openDynamicDataConfig'] = {'dynamicDataSourceConfigs': dynamic_data_source_configs}
|
||||
|
||||
url = f'{DINGTALK_OPENAPI_BASE}/v1.0/card/instances/createAndDeliver'
|
||||
headers = {
|
||||
'x-acs-dingtalk-access-token': self.access_token,
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
try:
|
||||
_stdout_logger.info(
|
||||
'DingTalk createAndDeliver request body: %s',
|
||||
json.dumps(body, ensure_ascii=False)[:1500],
|
||||
)
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(url, headers=headers, json=body, timeout=30.0)
|
||||
if response.status_code == 200:
|
||||
_stdout_logger.info(
|
||||
'DingTalk createAndDeliver response: %s',
|
||||
response.text[:500],
|
||||
)
|
||||
return True
|
||||
_stdout_logger.error(
|
||||
'DingTalk createAndDeliver failed: status=%s body=%s',
|
||||
response.status_code,
|
||||
response.text,
|
||||
)
|
||||
if self.logger:
|
||||
await self.logger.error(
|
||||
f'DingTalk createAndDeliver failed: status={response.status_code} body={response.text}'
|
||||
)
|
||||
return False
|
||||
except Exception:
|
||||
_stdout_logger.exception('DingTalk createAndDeliver error')
|
||||
if self.logger:
|
||||
await self.logger.error(f'DingTalk createAndDeliver error: {traceback.format_exc()}')
|
||||
return False
|
||||
|
||||
async def streaming_update_card(
|
||||
self,
|
||||
*,
|
||||
out_track_id: str,
|
||||
content_key: str,
|
||||
content_value: str,
|
||||
append: bool,
|
||||
finished: bool,
|
||||
failed: bool = False,
|
||||
) -> bool:
|
||||
"""PUT /v1.0/card/streaming.
|
||||
|
||||
Replaces `dingtalk_stream.AICardReplier.async_streaming` — same body
|
||||
shape (outTrackId / guid / key / content / isFull / isFinalize /
|
||||
isError) per the SDK source.
|
||||
"""
|
||||
if not await self.check_access_token():
|
||||
await self.get_access_token()
|
||||
|
||||
body = {
|
||||
'outTrackId': out_track_id,
|
||||
'guid': uuid.uuid4().hex,
|
||||
'key': content_key,
|
||||
'content': content_value,
|
||||
'isFull': not append,
|
||||
'isFinalize': finished,
|
||||
'isError': failed,
|
||||
}
|
||||
url = f'{DINGTALK_OPENAPI_BASE}/v1.0/card/streaming'
|
||||
headers = {
|
||||
'x-acs-dingtalk-access-token': self.access_token,
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.put(url, headers=headers, json=body, timeout=30.0)
|
||||
if response.status_code == 200:
|
||||
return True
|
||||
if self.logger:
|
||||
await self.logger.error(
|
||||
f'DingTalk card streaming failed: status={response.status_code} body={response.text}'
|
||||
)
|
||||
return False
|
||||
except Exception:
|
||||
if self.logger:
|
||||
await self.logger.error(f'DingTalk card streaming error: {traceback.format_exc()}')
|
||||
return False
|
||||
|
||||
async def update_card_data(
|
||||
self,
|
||||
*,
|
||||
out_track_id: str,
|
||||
card_param_map: Optional[dict] = None,
|
||||
private_data: Optional[dict] = None,
|
||||
) -> bool:
|
||||
"""PUT /v1.0/card/instances — non-streaming card content update."""
|
||||
if not await self.check_access_token():
|
||||
await self.get_access_token()
|
||||
|
||||
body: dict = {
|
||||
'outTrackId': out_track_id,
|
||||
'cardData': {'cardParamMap': _stringify_card_param_map(card_param_map)},
|
||||
}
|
||||
if private_data:
|
||||
body['privateData'] = private_data
|
||||
|
||||
url = f'{DINGTALK_OPENAPI_BASE}/v1.0/card/instances'
|
||||
headers = {
|
||||
'x-acs-dingtalk-access-token': self.access_token,
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
try:
|
||||
_stdout_logger.info(
|
||||
'DingTalk update_card_data request: out_track_id=%s body=%s',
|
||||
out_track_id,
|
||||
json.dumps(body, ensure_ascii=False)[:1500],
|
||||
)
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.put(url, headers=headers, json=body, timeout=30.0)
|
||||
_stdout_logger.info(
|
||||
'DingTalk update_card_data response: status=%d body=%s',
|
||||
response.status_code,
|
||||
response.text[:300],
|
||||
)
|
||||
if response.status_code == 200:
|
||||
return True
|
||||
if self.logger:
|
||||
await self.logger.error(
|
||||
f'DingTalk update card failed: status={response.status_code} body={response.text}'
|
||||
)
|
||||
return False
|
||||
except Exception:
|
||||
_stdout_logger.exception('DingTalk update_card_data error')
|
||||
if self.logger:
|
||||
await self.logger.error(f'DingTalk update card error: {traceback.format_exc()}')
|
||||
return False
|
||||
|
||||
async def get_legacy_access_token(self) -> Optional[str]:
|
||||
"""Fetch the LEGACY (oapi.dingtalk.com) access_token. This is a
|
||||
different auth domain from the v1.0 token cached in
|
||||
``self.access_token`` — only the legacy token authorises the
|
||||
``/media/upload`` endpoint that returns an ``@xxx`` media_id
|
||||
consumable by card components like Avatar.imageUrl.
|
||||
|
||||
Returns the token string on success, None on failure. Caches
|
||||
with a 60s safety margin before the documented 7200s expiry.
|
||||
"""
|
||||
now = time.time()
|
||||
if (
|
||||
self.legacy_access_token
|
||||
and self.legacy_access_token_expiry_time
|
||||
and now < self.legacy_access_token_expiry_time
|
||||
):
|
||||
return self.legacy_access_token
|
||||
|
||||
url = 'https://oapi.dingtalk.com/gettoken'
|
||||
try:
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.get(url, params={'appkey': self.key, 'appsecret': self.secret}, timeout=15.0)
|
||||
data = response.json() if response.status_code == 200 else {}
|
||||
if data.get('errcode') == 0 and data.get('access_token'):
|
||||
self.legacy_access_token = data['access_token']
|
||||
expires_in = int(data.get('expires_in', 7200))
|
||||
self.legacy_access_token_expiry_time = now + expires_in - 60
|
||||
return self.legacy_access_token
|
||||
if self.logger:
|
||||
await self.logger.error(
|
||||
f'DingTalk legacy gettoken failed: status={response.status_code} body={response.text[:200]}'
|
||||
)
|
||||
except Exception:
|
||||
_stdout_logger.exception('DingTalk legacy gettoken error')
|
||||
if self.logger:
|
||||
await self.logger.error(f'DingTalk legacy gettoken error: {traceback.format_exc()}')
|
||||
return None
|
||||
|
||||
async def upload_image_media(self, file_path: str) -> Optional[str]:
|
||||
"""Upload an image file to DingTalk media storage and return the
|
||||
``@xxx`` media_id, which can be passed straight into card variables
|
||||
like Avatar.imageUrl. Endpoint:
|
||||
|
||||
POST https://oapi.dingtalk.com/media/upload?access_token=…&type=image
|
||||
|
||||
Returns the media_id on success, None on any failure (caller
|
||||
should handle a None gracefully — DingTalk falls back to a
|
||||
default avatar when imageUrl is empty/unknown).
|
||||
"""
|
||||
if not os.path.exists(file_path):
|
||||
if self.logger:
|
||||
await self.logger.error(f'DingTalk upload_image_media: file not found {file_path}')
|
||||
return None
|
||||
|
||||
token = await self.get_legacy_access_token()
|
||||
if not token:
|
||||
return None
|
||||
|
||||
url = 'https://oapi.dingtalk.com/media/upload'
|
||||
try:
|
||||
with open(file_path, 'rb') as f:
|
||||
file_bytes = f.read()
|
||||
file_name = os.path.basename(file_path)
|
||||
# Best-effort content-type guess; DingTalk accepts the major image
|
||||
# mime types and otherwise infers from the bytes.
|
||||
ext = os.path.splitext(file_name)[1].lower().lstrip('.')
|
||||
mime = {'png': 'image/png', 'jpg': 'image/jpeg', 'jpeg': 'image/jpeg', 'gif': 'image/gif'}.get(
|
||||
ext, 'application/octet-stream'
|
||||
)
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(
|
||||
url,
|
||||
params={'access_token': token, 'type': 'image'},
|
||||
files={'media': (file_name, file_bytes, mime)},
|
||||
timeout=30.0,
|
||||
)
|
||||
data = response.json() if response.status_code == 200 else {}
|
||||
if data.get('errcode') == 0 and data.get('media_id'):
|
||||
_stdout_logger.info('DingTalk upload_image_media OK: media_id=%s', data['media_id'])
|
||||
return data['media_id']
|
||||
if self.logger:
|
||||
await self.logger.error(
|
||||
f'DingTalk upload_image_media failed: status={response.status_code} body={response.text[:300]}'
|
||||
)
|
||||
except Exception:
|
||||
_stdout_logger.exception('DingTalk upload_image_media error')
|
||||
if self.logger:
|
||||
await self.logger.error(f'DingTalk upload_image_media error: {traceback.format_exc()}')
|
||||
return None
|
||||
|
||||
async def start(self):
|
||||
"""启动 WebSocket 连接,监听消息"""
|
||||
self._stopped = False
|
||||
|
||||
@@ -1,106 +0,0 @@
|
||||
"""STREAM-mode handler for DingTalk card action button clicks.
|
||||
|
||||
DingTalk delivers card-action callbacks over the same WebSocket stream used
|
||||
for chatbot messages, under the topic `/v1.0/card/instances/callback`. This
|
||||
module subclasses `dingtalk_stream.CallbackHandler` and forwards the parsed
|
||||
payload to a coroutine the adapter registers, so the resume-paused-workflow
|
||||
logic stays in the platform adapter where it belongs.
|
||||
|
||||
The `CardCallbackMessage` returned by `from_dict` exposes:
|
||||
|
||||
* `card_instance_id` (from `outTrackId`) — the card whose button was clicked
|
||||
* `user_id` — the clicker's userId
|
||||
* `content` — parsed JSON; the click params live here. Where exactly inside
|
||||
`content` they sit depends on the template binding. We probe
|
||||
the common paths.
|
||||
* `extension` — parsed JSON; any extra data we set when delivering the card.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Awaitable, Callable, Optional
|
||||
|
||||
import dingtalk_stream # type: ignore
|
||||
from dingtalk_stream import AckMessage
|
||||
from dingtalk_stream.card_callback import CardCallbackMessage
|
||||
|
||||
|
||||
_PARAM_PATHS = (
|
||||
('params',),
|
||||
('cardPrivateData', 'params'),
|
||||
('userPrivateData', 'params'),
|
||||
('actionData', 'cardPrivateData', 'params'),
|
||||
)
|
||||
|
||||
|
||||
def _extract_params(content: dict) -> dict:
|
||||
"""Return the action params dict regardless of where the template put it."""
|
||||
for path in _PARAM_PATHS:
|
||||
node = content
|
||||
for key in path:
|
||||
if not isinstance(node, dict):
|
||||
node = None
|
||||
break
|
||||
node = node.get(key)
|
||||
if node is None:
|
||||
break
|
||||
if isinstance(node, dict) and node:
|
||||
return node
|
||||
return {}
|
||||
|
||||
|
||||
def _merge_params(*sources: dict) -> dict:
|
||||
merged = {}
|
||||
for source in sources:
|
||||
if isinstance(source, dict):
|
||||
merged.update(source)
|
||||
return merged
|
||||
|
||||
|
||||
class DingTalkCardActionHandler(dingtalk_stream.CallbackHandler):
|
||||
def __init__(
|
||||
self,
|
||||
dingtalk_stream_client,
|
||||
on_action: Optional[Callable[[dict], Awaitable[None]]] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dingtalk_client = dingtalk_stream_client
|
||||
self.on_action = on_action
|
||||
|
||||
async def process(self, callback: dingtalk_stream.CallbackMessage):
|
||||
try:
|
||||
message = CardCallbackMessage.from_dict(callback.data)
|
||||
content = message.content if isinstance(message.content, dict) else {}
|
||||
|
||||
# `CardCallbackMessage.from_dict` does not surface `actionId` (the
|
||||
# top-level field that ButtonGroup's sendCardRequest event puts
|
||||
# there). Pull it from the raw callback.data instead.
|
||||
raw = callback.data if isinstance(callback.data, dict) else {}
|
||||
params = _merge_params(_extract_params(content), _extract_params(raw))
|
||||
action_id = raw.get('actionId') or ''
|
||||
if not action_id:
|
||||
# Some templates nest it under actionData / cardPrivateData.
|
||||
action_data = raw.get('actionData') or {}
|
||||
if isinstance(action_data, dict):
|
||||
action_id = action_data.get('actionId') or action_id
|
||||
if not action_id:
|
||||
cpd = action_data.get('cardPrivateData') or {}
|
||||
if isinstance(cpd, dict):
|
||||
ids = cpd.get('actionIds')
|
||||
if isinstance(ids, list) and ids:
|
||||
action_id = str(ids[0])
|
||||
|
||||
payload = {
|
||||
'out_track_id': message.card_instance_id,
|
||||
'user_id': message.user_id,
|
||||
'corp_id': message.corp_id,
|
||||
'action_id': action_id,
|
||||
'params': params,
|
||||
'raw_content': message.content,
|
||||
'extension': message.extension if isinstance(message.extension, dict) else {},
|
||||
}
|
||||
if self.on_action is not None:
|
||||
await self.on_action(payload)
|
||||
except Exception as e:
|
||||
self.logger.error(f'DingTalkCardActionHandler.process error: {e}')
|
||||
return AckMessage.STATUS_OK, 'OK'
|
||||
@@ -12,142 +12,6 @@ import traceback
|
||||
from cryptography.hazmat.primitives.asymmetric import ed25519
|
||||
|
||||
|
||||
QQ_SELECT_ACTION_PREFIX = '__langbot_select__:'
|
||||
|
||||
|
||||
def get_select_field_options(form_data: dict) -> tuple[str, list[str]]:
|
||||
"""Return the active select field name and its display/submission values."""
|
||||
field_name = str(form_data.get('_current_input_field') or '').strip()
|
||||
if not field_name:
|
||||
return '', []
|
||||
|
||||
field = next(
|
||||
(
|
||||
item
|
||||
for item in form_data.get('input_defs') or []
|
||||
if str(item.get('output_variable_name') or '').strip() == field_name
|
||||
),
|
||||
None,
|
||||
)
|
||||
if not field or str(field.get('type') or '').strip().lower() != 'select':
|
||||
return '', []
|
||||
|
||||
source = field.get('option_source') or {}
|
||||
source_value = source.get('value') if isinstance(source, dict) else None
|
||||
if isinstance(source_value, list):
|
||||
return field_name, [str(item) for item in source_value]
|
||||
if isinstance(source_value, str):
|
||||
return field_name, [part.strip() for part in source_value.splitlines() if part.strip()]
|
||||
|
||||
options = field.get('options')
|
||||
if not isinstance(options, list):
|
||||
return field_name, []
|
||||
values = []
|
||||
for item in options:
|
||||
if isinstance(item, dict):
|
||||
values.append(str(item.get('label') or item.get('value') or ''))
|
||||
else:
|
||||
values.append(str(item))
|
||||
return field_name, [value for value in values if value]
|
||||
|
||||
|
||||
def build_keyboard_from_select_field(form_data: dict, *, buttons_per_row: int | None = None) -> dict:
|
||||
"""Build callback buttons for the currently active Dify select field."""
|
||||
_, options = get_select_field_options(form_data)
|
||||
visible_options = options[:25]
|
||||
if buttons_per_row is None:
|
||||
# Keep small choices readable while fitting up to QQ's 5x5 limit.
|
||||
buttons_per_row = min(5, max(2, (len(visible_options) + 4) // 5))
|
||||
selection_actions = [
|
||||
{
|
||||
'id': f'{QQ_SELECT_ACTION_PREFIX}{idx}',
|
||||
'title': option,
|
||||
'button_style': 'secondary',
|
||||
}
|
||||
for idx, option in enumerate(visible_options)
|
||||
]
|
||||
return build_keyboard_from_form({'actions': selection_actions}, buttons_per_row=buttons_per_row)
|
||||
|
||||
|
||||
def resolve_select_button_action(form_data: dict, action_id: str) -> tuple[str, str] | None:
|
||||
"""Resolve a select-button callback to ``(field_name, option_value)``."""
|
||||
if not action_id.startswith(QQ_SELECT_ACTION_PREFIX):
|
||||
return None
|
||||
try:
|
||||
option_index = int(action_id[len(QQ_SELECT_ACTION_PREFIX) :])
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
field_name, options = get_select_field_options(form_data)
|
||||
if not field_name or option_index < 0 or option_index >= len(options) or option_index >= 25:
|
||||
return None
|
||||
return field_name, options[option_index]
|
||||
|
||||
|
||||
def build_keyboard_from_form(form_data: dict, *, buttons_per_row: int = 2) -> dict:
|
||||
"""Build a QQ keyboard JSON payload from a Dify human-input form_data.
|
||||
|
||||
Each Dify ``action`` becomes a callback button (``action.type=1``)
|
||||
whose ``data`` is set directly to the Dify ``action_id``. The
|
||||
INTERACTION_CREATE event carries this back as
|
||||
``data.resolved.button_data`` so the adapter can match the click to
|
||||
the originating form.
|
||||
|
||||
Layout limits per spec: max 5 rows, max 5 buttons per row. We default
|
||||
to 2 buttons per row for legibility; oversized button lists wrap
|
||||
onto additional rows and overflow gets dropped (max 25 visible).
|
||||
|
||||
Args:
|
||||
form_data: Dify ``{"actions": [{"id", "title", "button_style"}, ...]}``.
|
||||
buttons_per_row: 1..5. Mobile UI looks best at 2.
|
||||
|
||||
Returns:
|
||||
``{"content": {"rows": [{"buttons": [...]}]}}``.
|
||||
"""
|
||||
actions = list(form_data.get('actions') or [])[:25] # 5×5 hard cap
|
||||
buttons_per_row = max(1, min(5, buttons_per_row))
|
||||
|
||||
def _button(idx: int, action: dict) -> dict:
|
||||
action_id = str(action.get('id') or '')
|
||||
label = str(action.get('title') or action_id or f'选项 {idx + 1}')
|
||||
style_raw = (action.get('button_style') or '').lower()
|
||||
# QQ: 0 灰色线框, 1 蓝色线框. Highlight the primary / first action.
|
||||
if style_raw == 'primary' or (style_raw == '' and idx == 0):
|
||||
style = 1
|
||||
else:
|
||||
style = 0
|
||||
return {
|
||||
'id': str(idx + 1),
|
||||
'render_data': {
|
||||
'label': label,
|
||||
# Shown after the user clicks — gives local "已选择" feedback
|
||||
# without a follow-up message. Style mimics DingTalk/Lark's
|
||||
# in-card selection state.
|
||||
'visited_label': f'✓ {label}',
|
||||
'style': style,
|
||||
},
|
||||
'action': {
|
||||
'type': 1, # callback button
|
||||
'permission': {'type': 2}, # everyone can click
|
||||
'data': action_id,
|
||||
'unsupport_tips': '当前客户端版本不支持此按钮,请升级 QQ',
|
||||
},
|
||||
}
|
||||
|
||||
rows = []
|
||||
for row_start in range(0, len(actions), buttons_per_row):
|
||||
row_actions = actions[row_start : row_start + buttons_per_row]
|
||||
rows.append(
|
||||
{
|
||||
'buttons': [_button(row_start + j, a) for j, a in enumerate(row_actions)],
|
||||
}
|
||||
)
|
||||
if len(rows) >= 5:
|
||||
break
|
||||
|
||||
return {'content': {'rows': rows}}
|
||||
|
||||
|
||||
class QQOfficialClient:
|
||||
def __init__(self, secret: str, token: str, app_id: str, logger: None, unified_mode: bool = False):
|
||||
self.unified_mode = unified_mode
|
||||
@@ -166,10 +30,6 @@ class QQOfficialClient:
|
||||
self.token = token
|
||||
self.app_id = app_id
|
||||
self._message_handlers = {}
|
||||
# Single optional handler for INTERACTION_CREATE (button click). We
|
||||
# don't multiplex like message handlers — only the adapter cares,
|
||||
# and the click<->resume path needs a single source of truth.
|
||||
self._interaction_handler: Optional[Callable[[Dict[str, Any], Optional[str]], Any]] = None
|
||||
self.base_url = 'https://api.sgroup.qq.com'
|
||||
self.access_token = ''
|
||||
self.access_token_expiry_time = None
|
||||
@@ -247,23 +107,6 @@ class QQOfficialClient:
|
||||
return response, 200
|
||||
|
||||
if payload.get('op') == 0:
|
||||
# INTERACTION_CREATE (button click) skips ``get_message`` —
|
||||
# that helper only flattens message-event fields and would
|
||||
# drop ``data.resolved.button_data`` / ``data.button_id``.
|
||||
if payload.get('t') == 'INTERACTION_CREATE':
|
||||
if self._interaction_handler:
|
||||
try:
|
||||
d = payload.get('d') or {}
|
||||
# Top-level ``id`` is the ws/event id used as
|
||||
# ``event_id`` for passive replies. ``d.id``
|
||||
# is the interaction id used for ACK. Do not
|
||||
# confuse the two — QQ rejects misuse with
|
||||
# 40034025.
|
||||
ws_event_id = payload.get('id')
|
||||
await self._interaction_handler(d, ws_event_id)
|
||||
except Exception:
|
||||
await self.logger.error(f'Error in interaction handler: {traceback.format_exc()}')
|
||||
return {'code': 0, 'message': 'success'}
|
||||
message_data = await self.get_message(payload)
|
||||
if message_data:
|
||||
event = QQOfficialEvent.from_payload(message_data)
|
||||
@@ -290,21 +133,6 @@ class QQOfficialClient:
|
||||
|
||||
return decorator
|
||||
|
||||
def on_interaction(self):
|
||||
"""Register a single handler for INTERACTION_CREATE events.
|
||||
|
||||
The handler receives ``(data_dict, interaction_id)`` — the raw
|
||||
``d`` payload plus the top-level ``id`` field (the interaction
|
||||
id, needed for the PUT /interactions/{id} ack and for reuse as
|
||||
an ``event_id`` on the resumed reply within 30 minutes).
|
||||
"""
|
||||
|
||||
def decorator(func: Callable[[Dict[str, Any], Optional[str]], Any]):
|
||||
self._interaction_handler = func
|
||||
return func
|
||||
|
||||
return decorator
|
||||
|
||||
async def _handle_message(self, event: QQOfficialEvent):
|
||||
"""处理消息事件"""
|
||||
msg_type = event.t
|
||||
@@ -349,20 +177,8 @@ class QQOfficialClient:
|
||||
content_type = attachment.get('content_type', '')
|
||||
return content_type.startswith('image/')
|
||||
|
||||
async def send_private_text_msg(
|
||||
self,
|
||||
user_openid: str,
|
||||
content: str,
|
||||
msg_id: Optional[str] = None,
|
||||
event_id: Optional[str] = None,
|
||||
msg_seq: int = 1,
|
||||
):
|
||||
"""Send a c2c text message.
|
||||
|
||||
Either ``msg_id`` (inbound user msg, free passive reply) or
|
||||
``event_id`` (e.g. INTERACTION_CREATE id, valid 30 min) is
|
||||
required. Without either, the call costs the proactive-send quota.
|
||||
"""
|
||||
async def send_private_text_msg(self, user_openid: str, content: str, msg_id: str):
|
||||
"""发送私聊消息"""
|
||||
if not await self.check_access_token():
|
||||
await self.get_access_token()
|
||||
|
||||
@@ -372,15 +188,11 @@ class QQOfficialClient:
|
||||
'Authorization': f'QQBot {self.access_token}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
data: dict[str, Any] = {
|
||||
data = {
|
||||
'content': content,
|
||||
'msg_type': 0,
|
||||
'msg_seq': msg_seq,
|
||||
'msg_id': msg_id,
|
||||
}
|
||||
if msg_id:
|
||||
data['msg_id'] = msg_id
|
||||
if event_id:
|
||||
data['event_id'] = event_id
|
||||
response = await client.post(url, headers=headers, json=data)
|
||||
response_data = response.json()
|
||||
if response.status_code == 200:
|
||||
@@ -389,19 +201,8 @@ class QQOfficialClient:
|
||||
await self.logger.error(f'Failed to send private message: {response_data}')
|
||||
raise ValueError(response)
|
||||
|
||||
async def send_group_text_msg(
|
||||
self,
|
||||
group_openid: str,
|
||||
content: str,
|
||||
msg_id: Optional[str] = None,
|
||||
event_id: Optional[str] = None,
|
||||
msg_seq: int = 1,
|
||||
):
|
||||
"""Send a group text message.
|
||||
|
||||
Either ``msg_id`` or ``event_id`` is required (see
|
||||
:meth:`send_private_text_msg` for the distinction).
|
||||
"""
|
||||
async def send_group_text_msg(self, group_openid: str, content: str, msg_id: str):
|
||||
"""发送群聊消息"""
|
||||
if not await self.check_access_token():
|
||||
await self.get_access_token()
|
||||
|
||||
@@ -411,15 +212,11 @@ class QQOfficialClient:
|
||||
'Authorization': f'QQBot {self.access_token}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
data: dict[str, Any] = {
|
||||
data = {
|
||||
'content': content,
|
||||
'msg_type': 0,
|
||||
'msg_seq': msg_seq,
|
||||
'msg_id': msg_id,
|
||||
}
|
||||
if msg_id:
|
||||
data['msg_id'] = msg_id
|
||||
if event_id:
|
||||
data['event_id'] = event_id
|
||||
response = await client.post(url, headers=headers, json=data)
|
||||
if response.status_code == 200:
|
||||
return
|
||||
@@ -688,107 +485,6 @@ class QQOfficialClient:
|
||||
raise Exception(f'Failed to send stream message: HTTP {response.status_code} {response.text}')
|
||||
return response.json()
|
||||
|
||||
async def send_markdown_keyboard(
|
||||
self,
|
||||
target_type: str,
|
||||
target_id: str,
|
||||
markdown_content: str,
|
||||
keyboard: Optional[dict] = None,
|
||||
msg_id: Optional[str] = None,
|
||||
event_id: Optional[str] = None,
|
||||
msg_seq: int = 1,
|
||||
) -> dict:
|
||||
"""Send a ``msg_type=2`` (markdown) message carrying a keyboard.
|
||||
|
||||
The keyboard ride-along is the only documented way to attach
|
||||
buttons in QQ official; pure keyboard-only messages are not
|
||||
accepted by the server (markdown content is required).
|
||||
|
||||
Args:
|
||||
target_type: 'c2c' (single chat), 'group', 'channel' (text
|
||||
channel — uses POST /channels/{id}/messages instead of v2).
|
||||
target_id: openid for c2c/group, channel_id for channel.
|
||||
markdown_content: Plain markdown text shown above the buttons.
|
||||
keyboard: ``{'content': {'rows': [{'buttons': [...]}]}}`` per
|
||||
the official spec. Use :func:`build_keyboard_from_form`
|
||||
to construct from Dify form_data.
|
||||
msg_id: Inbound user message id; turns this into a passive
|
||||
reply (preferred — no monthly quota cost).
|
||||
event_id: Use ``INTERACTION_CREATE`` event id from a prior
|
||||
button click to keep within the 30-minute passive window
|
||||
without an inbound msg_id.
|
||||
msg_seq: De-dup counter when reusing msg_id.
|
||||
"""
|
||||
if not await self.check_access_token():
|
||||
await self.get_access_token()
|
||||
|
||||
if target_type == 'c2c':
|
||||
url = f'{self.base_url}/v2/users/{target_id}/messages'
|
||||
elif target_type == 'group':
|
||||
url = f'{self.base_url}/v2/groups/{target_id}/messages'
|
||||
elif target_type == 'channel':
|
||||
url = f'{self.base_url}/channels/{target_id}/messages'
|
||||
else:
|
||||
raise ValueError(f'Unsupported target_type for markdown+keyboard: {target_type}')
|
||||
|
||||
body: dict[str, Any] = {
|
||||
'msg_type': 2,
|
||||
'markdown': {'content': markdown_content},
|
||||
'msg_seq': msg_seq,
|
||||
}
|
||||
if keyboard and keyboard.get('content', {}).get('rows'):
|
||||
body['keyboard'] = keyboard
|
||||
if msg_id:
|
||||
body['msg_id'] = msg_id
|
||||
if event_id:
|
||||
body['event_id'] = event_id
|
||||
|
||||
async with httpx.AsyncClient(timeout=30) as client:
|
||||
headers = {
|
||||
'Authorization': f'QQBot {self.access_token}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
response = await client.post(url, headers=headers, json=body)
|
||||
if response.status_code != 200:
|
||||
await self.logger.error(
|
||||
f'Failed to send markdown+keyboard: HTTP {response.status_code} {response.text}'
|
||||
)
|
||||
raise Exception(f'Failed to send markdown+keyboard: HTTP {response.status_code} {response.text}')
|
||||
return response.json()
|
||||
|
||||
async def ack_interaction(self, interaction_id: str, code: int = 0) -> None:
|
||||
"""Acknowledge a button-click INTERACTION_CREATE event.
|
||||
|
||||
QQ keeps the client in a loading spinner until this ack is
|
||||
received. Should be called as soon as the click is parsed, before
|
||||
any heavier downstream work (the actual workflow resume can run
|
||||
async).
|
||||
|
||||
Args:
|
||||
interaction_id: The ``id`` field from the INTERACTION_CREATE event.
|
||||
code: 0=success, 1=fail, 2=rate-limited, 3=duplicate, 4=no
|
||||
permission, 5=admin only. Default 0.
|
||||
"""
|
||||
if not interaction_id:
|
||||
return
|
||||
if not await self.check_access_token():
|
||||
await self.get_access_token()
|
||||
|
||||
url = f'{self.base_url}/interactions/{interaction_id}'
|
||||
async with httpx.AsyncClient(timeout=10) as client:
|
||||
headers = {
|
||||
'Authorization': f'QQBot {self.access_token}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
try:
|
||||
response = await client.put(url, headers=headers, json={'code': code})
|
||||
if response.status_code >= 400:
|
||||
await self.logger.warning(
|
||||
f'ack_interaction non-success: HTTP {response.status_code} {response.text}'
|
||||
)
|
||||
except Exception as e:
|
||||
await self.logger.warning(f'ack_interaction error (non-fatal): {e}')
|
||||
|
||||
async def is_token_expired(self):
|
||||
"""检查token是否过期"""
|
||||
if self.access_token_expiry_time is None:
|
||||
@@ -957,12 +653,6 @@ class QQOfficialClient:
|
||||
d = payload.get('d', {})
|
||||
s = payload.get('s')
|
||||
t = payload.get('t')
|
||||
# Top-level event id, distinct from `d.id`. Per QQ
|
||||
# spec this is the only value accepted as ``event_id``
|
||||
# in subsequent passive-reply send-message calls
|
||||
# (``d.id`` for INTERACTION_CREATE is the interaction
|
||||
# id, used solely for PUT /interactions/{id} ack).
|
||||
ws_event_id = payload.get('id')
|
||||
|
||||
if not isinstance(d, dict):
|
||||
d = {}
|
||||
@@ -1041,22 +731,7 @@ class QQOfficialClient:
|
||||
|
||||
else:
|
||||
await self.logger.debug(f'Received event: {t}, seq={s}')
|
||||
# INTERACTION_CREATE bypasses the regular
|
||||
# on_event dispatcher so the adapter sees the
|
||||
# top-level ws_event_id (needed as event_id
|
||||
# for the resumed reply) — same shape as the
|
||||
# webhook handler.
|
||||
if t == 'INTERACTION_CREATE':
|
||||
if self._interaction_handler:
|
||||
try:
|
||||
result = self._interaction_handler(d, ws_event_id)
|
||||
if asyncio.iscoroutine(result):
|
||||
await result
|
||||
except Exception:
|
||||
await self.logger.error(
|
||||
f'Error in interaction handler (ws): {traceback.format_exc()}'
|
||||
)
|
||||
elif on_event:
|
||||
if on_event:
|
||||
try:
|
||||
result = on_event(t, d)
|
||||
if asyncio.iscoroutine(result):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -20,19 +20,7 @@ from typing import Any, Callable, Optional
|
||||
import aiohttp
|
||||
|
||||
from langbot.libs.wecom_ai_bot_api import wecombotevent
|
||||
from langbot.libs.wecom_ai_bot_api.api import (
|
||||
parse_wecom_bot_message,
|
||||
StreamSession,
|
||||
build_human_input_template_card_payload,
|
||||
build_human_input_text_prompt,
|
||||
build_button_interaction_update_card,
|
||||
build_multiple_interaction_update_card,
|
||||
extract_template_card_action,
|
||||
extract_template_card_event_payload,
|
||||
extract_template_card_selections,
|
||||
extract_wecom_event_type,
|
||||
parse_select_button_action,
|
||||
)
|
||||
from langbot.libs.wecom_ai_bot_api.api import parse_wecom_bot_message, StreamSession
|
||||
from langbot.pkg.platform.logger import EventLogger
|
||||
|
||||
DEFAULT_WS_URL = 'wss://openws.work.weixin.qq.com'
|
||||
@@ -55,10 +43,6 @@ def _generate_req_id(prefix: str) -> str:
|
||||
return f'{prefix}_{ts}_{rand}'
|
||||
|
||||
|
||||
def _frame_snippet(frame: dict, limit: int = 1000) -> str:
|
||||
return json.dumps(frame, ensure_ascii=False, default=str)[:limit]
|
||||
|
||||
|
||||
class WecomBotWsClient:
|
||||
"""WeChat Work AI Bot WebSocket long connection client.
|
||||
|
||||
@@ -119,22 +103,6 @@ class WecomBotWsClient:
|
||||
# msg_id -> feedback_id (for associating feedback with message)
|
||||
self._msg_feedback_ids: dict[str, str] = {} # msg_id -> feedback_id
|
||||
|
||||
# Dify human-input pause state for ws mode. Keys are task_id (echoed
|
||||
# back in template_card_event.TaskId so we can rebuild the session
|
||||
# context on click).
|
||||
# task_id -> {form_data, msg_id, user_id, chat_id, stream_id, req_id}
|
||||
self._pending_forms_by_task: dict[str, dict] = {}
|
||||
# Reverse: msg_id -> task_id (for cleanup when stream finishes).
|
||||
self._task_id_by_msg: dict[str, str] = {}
|
||||
# Optional card-action callback registered by the adapter.
|
||||
# Signature mirrors the http-mode WecomBotClient:
|
||||
# async def callback(session, action_id, task_id, raw_event) -> None
|
||||
self._card_action_callback: Optional[Callable] = None
|
||||
# Optional `source` block injected into every interactive
|
||||
# template_card the client builds via `push_form_pause`. Set via
|
||||
# `set_card_source` from the adapter after reading config.
|
||||
self.card_source: Optional[dict] = None
|
||||
|
||||
# ── Public API ──────────────────────────────────────────────────
|
||||
|
||||
async def connect(self):
|
||||
@@ -268,132 +236,6 @@ class WecomBotWsClient:
|
||||
}
|
||||
return await self._send_reply(req_id, body)
|
||||
|
||||
async def reply_template_card(self, req_id: str, card_payload: dict[str, Any]) -> Optional[dict]:
|
||||
"""Send a template_card (button_interaction etc.) reply.
|
||||
|
||||
Args:
|
||||
req_id: The req_id from the original message frame.
|
||||
card_payload: Body produced by ``build_button_interaction_payload``;
|
||||
must contain ``msgtype`` and ``template_card`` keys.
|
||||
|
||||
Returns:
|
||||
ACK frame dict, or None on failure.
|
||||
"""
|
||||
return await self._send_reply(req_id, card_payload)
|
||||
|
||||
async def update_template_card(
|
||||
self,
|
||||
req_id: str,
|
||||
template_card: dict[str, Any],
|
||||
) -> Optional[dict]:
|
||||
"""Update an existing template_card via WebSocket.
|
||||
|
||||
Uses the ``aibot_respond_update_msg`` command. Must be called
|
||||
within 5 seconds of receiving the ``template_card_event`` callback,
|
||||
using the **same req_id** from that callback.
|
||||
|
||||
The ``template_card`` dict should contain ``card_type`` and the
|
||||
new content fields (e.g. ``main_title``, ``button_list`` with
|
||||
disabled buttons and ``replace_text``).
|
||||
|
||||
Returns:
|
||||
ACK frame dict, or None on failure.
|
||||
"""
|
||||
body: dict[str, Any] = {
|
||||
'response_type': 'update_template_card',
|
||||
'template_card': template_card,
|
||||
}
|
||||
return await self._send_reply(req_id, body, cmd=CMD_RESPOND_UPDATE)
|
||||
|
||||
def set_card_action_callback(self, callback: Callable) -> None:
|
||||
"""Register the button-click handler.
|
||||
|
||||
``async def callback(session, action_id, task_id, raw_event) -> None``
|
||||
— same signature as the http-mode WecomBotClient version so the
|
||||
adapter can hand both off to the same coroutine.
|
||||
"""
|
||||
self._card_action_callback = callback
|
||||
|
||||
def set_card_source(self, source: Optional[dict]) -> None:
|
||||
"""Set the `source` block injected into every interactive
|
||||
template_card pushed via `push_form_pause`. Pass None to clear."""
|
||||
self.card_source = source
|
||||
|
||||
async def push_form_pause(
|
||||
self, msg_id: str, form_data: dict, task_id: Optional[str] = None
|
||||
) -> tuple[bool, Optional[str], Optional[str]]:
|
||||
"""Attach a Dify human-input pause to the active stream and send
|
||||
the button_interaction card immediately.
|
||||
|
||||
ws mode has no notion of polled "followup" responses — each reply
|
||||
is a one-shot frame send. So unlike the http path (which defers
|
||||
card delivery to the next followup), here we just craft the card
|
||||
and reply with it on the original req_id. The corresponding stream
|
||||
session is then torn down so subsequent chunks don't re-send.
|
||||
|
||||
Returns:
|
||||
``(ok, stream_id, task_id)``. ``ok=False`` if no active stream
|
||||
for this msg_id (e.g. message arrived in non-stream mode).
|
||||
"""
|
||||
key = self._stream_ids.get(msg_id)
|
||||
if not key:
|
||||
return False, None, None
|
||||
req_id, stream_id = key.split('|', 1)
|
||||
|
||||
if not task_id:
|
||||
task_id = f'dify-{secrets.token_hex(12)}'
|
||||
|
||||
session_info = self._stream_sessions.get(msg_id) or {}
|
||||
text_prompt = build_human_input_text_prompt(form_data)
|
||||
if text_prompt:
|
||||
try:
|
||||
ack = await self.reply_text(req_id, text_prompt)
|
||||
if ack is None:
|
||||
return False, stream_id, None
|
||||
except Exception:
|
||||
await self.logger.error(f'Failed to send human-input text prompt: {traceback.format_exc()}')
|
||||
return False, stream_id, None
|
||||
|
||||
self._stream_ids.pop(msg_id, None)
|
||||
self._stream_last_content.pop(msg_id, None)
|
||||
self._stream_sessions.pop(msg_id, None)
|
||||
return True, stream_id, None
|
||||
|
||||
self._pending_forms_by_task[task_id] = {
|
||||
'form_data': form_data,
|
||||
'msg_id': msg_id,
|
||||
'user_id': session_info.get('user_id', ''),
|
||||
'chat_id': session_info.get('chat_id', ''),
|
||||
'stream_id': stream_id,
|
||||
'req_id': req_id,
|
||||
}
|
||||
self._task_id_by_msg[msg_id] = task_id
|
||||
|
||||
card_payload = build_human_input_template_card_payload(
|
||||
form_data,
|
||||
task_id,
|
||||
source=self.card_source,
|
||||
select_as_buttons=True,
|
||||
)
|
||||
try:
|
||||
await self.reply_template_card(req_id, card_payload)
|
||||
except Exception:
|
||||
await self.logger.error(f'Failed to send button_interaction card: {traceback.format_exc()}')
|
||||
# Roll back the bookkeeping so the next attempt isn't blocked.
|
||||
self._pending_forms_by_task.pop(task_id, None)
|
||||
self._task_id_by_msg.pop(msg_id, None)
|
||||
return False, stream_id, None
|
||||
|
||||
# Tear down the stream — WeCom expects either stream chunks OR a
|
||||
# template_card, not both on the same req_id. Subsequent
|
||||
# push_stream_chunk calls for this msg_id become no-ops.
|
||||
self._stream_ids.pop(msg_id, None)
|
||||
self._stream_last_content.pop(msg_id, None)
|
||||
# Keep _stream_sessions so the button callback can still resolve
|
||||
# user/chat context; it gets cleaned up when the click fires.
|
||||
|
||||
return True, stream_id, task_id
|
||||
|
||||
async def send_message(self, chat_id: str, content: str, msgtype: str = 'markdown') -> Optional[dict]:
|
||||
"""Proactively send a message to a specified chat.
|
||||
|
||||
@@ -416,23 +258,6 @@ class WecomBotWsClient:
|
||||
body['text'] = {'content': content}
|
||||
return await self._send_reply(req_id, body, cmd=CMD_SEND_MSG)
|
||||
|
||||
async def send_template_card(self, chat_id: str, card_payload: dict[str, Any]) -> Optional[dict]:
|
||||
"""Proactively push a template_card to a chat.
|
||||
|
||||
Used for the resumed-workflow path (button click → new query):
|
||||
synthetic events have no inbound req_id to reply against, so we
|
||||
fall back to proactive ``aibot_send_msg`` instead of reply mode.
|
||||
|
||||
Args:
|
||||
chat_id: userid (single chat) or chatid (group chat).
|
||||
card_payload: ``{"msgtype": "template_card", "template_card": {...}}``
|
||||
as produced by :func:`build_button_interaction_payload`.
|
||||
"""
|
||||
req_id = _generate_req_id(CMD_SEND_MSG)
|
||||
body = dict(card_payload)
|
||||
body['chatid'] = chat_id
|
||||
return await self._send_reply(req_id, body, cmd=CMD_SEND_MSG)
|
||||
|
||||
async def push_stream_chunk(self, msg_id: str, content: str, is_final: bool = False) -> bool:
|
||||
"""Push a streaming chunk for a given message ID.
|
||||
|
||||
@@ -451,31 +276,10 @@ class WecomBotWsClient:
|
||||
return False
|
||||
req_id, stream_id = key.split('|', 1)
|
||||
try:
|
||||
previous_content = self._stream_last_content.get(msg_id, '')
|
||||
if previous_content and content.startswith(previous_content):
|
||||
next_content = content
|
||||
elif previous_content and not content:
|
||||
next_content = previous_content
|
||||
else:
|
||||
next_content = previous_content + content if previous_content else content
|
||||
|
||||
# Skip sending if content hasn't changed (e.g. during tool call argument streaming)
|
||||
if not is_final and next_content == previous_content:
|
||||
if not is_final and content == self._stream_last_content.get(msg_id):
|
||||
return True
|
||||
|
||||
# Skip empty/whitespace-only snapshots — the runner injects a
|
||||
# zero-width space ('') as a pass-through when workflow_paused
|
||||
# fires without any preceding LLM output. WeCom renders that
|
||||
# as an empty bubble that sits before the form card; skip it.
|
||||
# NOTE: Python str.strip() does NOT strip , so we use
|
||||
# a regex that treats any character with Unicode category Zs
|
||||
# (separator space) or Cf (format char like ZWS) as blank.
|
||||
if not is_final:
|
||||
import re as _re
|
||||
|
||||
if not _re.sub(r'[\s]', '', next_content):
|
||||
return True
|
||||
|
||||
# Generate feedback_id for final chunk
|
||||
feedback_id = ''
|
||||
if is_final:
|
||||
@@ -486,10 +290,8 @@ class WecomBotWsClient:
|
||||
if session_info:
|
||||
self._feedback_sessions[feedback_id] = session_info
|
||||
|
||||
# WeCom replaces the displayed stream content on each refresh, so
|
||||
# every frame must contain the complete snapshot, not only a delta.
|
||||
await self.reply_stream(req_id, stream_id, next_content, finish=is_final, feedback_id=feedback_id)
|
||||
self._stream_last_content[msg_id] = next_content
|
||||
await self.reply_stream(req_id, stream_id, content, finish=is_final, feedback_id=feedback_id)
|
||||
self._stream_last_content[msg_id] = content
|
||||
if is_final:
|
||||
self._stream_ids.pop(msg_id, None)
|
||||
self._stream_last_content.pop(msg_id, None)
|
||||
@@ -663,7 +465,7 @@ class WecomBotWsClient:
|
||||
return
|
||||
|
||||
# Unknown frame
|
||||
await self.logger.warning(f'Unknown frame: {_frame_snippet(frame)}')
|
||||
await self.logger.warning(f'Unknown frame: {json.dumps(frame, ensure_ascii=False)[:200]}')
|
||||
|
||||
async def _handle_message_callback(self, frame: dict):
|
||||
"""Handle an incoming message callback frame."""
|
||||
@@ -671,13 +473,6 @@ class WecomBotWsClient:
|
||||
body = frame.get('body', {})
|
||||
req_id = frame.get('headers', {}).get('req_id', '')
|
||||
|
||||
event_type = extract_wecom_event_type(body)
|
||||
if event_type == 'template_card_event':
|
||||
await self._handle_template_card_event_frame(frame, body)
|
||||
return
|
||||
if event_type:
|
||||
await self.logger.debug(f'Received msg_callback event_type={event_type}: {_frame_snippet(frame)}')
|
||||
|
||||
# Parse message using shared logic
|
||||
message_data = await parse_wecom_bot_message(body, self.encoding_aes_key, self.logger)
|
||||
if not message_data:
|
||||
@@ -711,12 +506,8 @@ class WecomBotWsClient:
|
||||
body = frame.get('body', {})
|
||||
req_id = frame.get('headers', {}).get('req_id', '')
|
||||
|
||||
event_info = body.get('event', {}) if isinstance(body.get('event'), dict) else body
|
||||
event_type = extract_wecom_event_type(body)
|
||||
if not event_type:
|
||||
await self.logger.warning(f'Received event_callback without event_type: {_frame_snippet(frame)}')
|
||||
else:
|
||||
await self.logger.debug(f'Received event_callback event_type={event_type}')
|
||||
event_info = body.get('event', {})
|
||||
event_type = event_info.get('eventtype', '')
|
||||
|
||||
message_data = {
|
||||
'msgtype': 'event',
|
||||
@@ -777,10 +568,6 @@ class WecomBotWsClient:
|
||||
await self.logger.error(f'Error in feedback handler: {traceback.format_exc()}')
|
||||
return
|
||||
|
||||
if event_type == 'template_card_event':
|
||||
await self._handle_template_card_event_frame(frame, body)
|
||||
return
|
||||
|
||||
event = wecombotevent.WecomBotEvent(message_data)
|
||||
|
||||
if event_type in self._message_handlers:
|
||||
@@ -794,72 +581,6 @@ class WecomBotWsClient:
|
||||
except Exception:
|
||||
await self.logger.error(f'Error in event callback: {traceback.format_exc()}')
|
||||
|
||||
async def _handle_template_card_event_frame(self, frame: dict, body: dict):
|
||||
"""Handle template_card_event frames from event_callback or msg_callback."""
|
||||
tce = extract_template_card_event_payload(body)
|
||||
task_id, event_key, card_type = extract_template_card_action(tce)
|
||||
await self.logger.info(
|
||||
f'Received template_card_event (ws): task_id={task_id} event_key={event_key!r} card_type={card_type}'
|
||||
)
|
||||
|
||||
pending = self._pending_forms_by_task.get(task_id)
|
||||
if pending is None:
|
||||
await self.logger.warning(f'No pending_form found for task_id={task_id} (ws); card event ignored')
|
||||
return
|
||||
|
||||
req_id_for_update = frame.get('headers', {}).get('req_id', '')
|
||||
form_data = pending.get('form_data', {}) or {}
|
||||
selections = extract_template_card_selections(tce, form_data)
|
||||
if not selections:
|
||||
selections = parse_select_button_action(event_key, form_data)
|
||||
if card_type == 'multiple_interaction' and not selections:
|
||||
await self.logger.warning(
|
||||
f'multiple_interaction callback has no parseable selections (ws): raw={str(tce)[:1000]}'
|
||||
)
|
||||
self._drop_pending_form_task(task_id, pending)
|
||||
return
|
||||
|
||||
update_card = build_button_interaction_update_card(
|
||||
form_data,
|
||||
task_id,
|
||||
event_key,
|
||||
source=self.card_source,
|
||||
)
|
||||
if card_type == 'multiple_interaction' or selections:
|
||||
update_card = build_multiple_interaction_update_card(
|
||||
form_data,
|
||||
task_id,
|
||||
selections,
|
||||
source=self.card_source,
|
||||
)
|
||||
try:
|
||||
await self.update_template_card(req_id_for_update, update_card)
|
||||
except Exception:
|
||||
await self.logger.warning(f'Failed to update template card (ws): {traceback.format_exc()}')
|
||||
|
||||
if self._card_action_callback is not None:
|
||||
try:
|
||||
session = StreamSession(
|
||||
stream_id=pending.get('stream_id', ''),
|
||||
msg_id=pending.get('msg_id', ''),
|
||||
chat_id=pending.get('chat_id') or None,
|
||||
user_id=pending.get('user_id') or None,
|
||||
)
|
||||
session.pending_form = pending.get('form_data')
|
||||
session.pending_form_task_id = task_id
|
||||
await self._card_action_callback(session, event_key, task_id, body)
|
||||
except Exception:
|
||||
await self.logger.error(f'card action callback raised (ws): {traceback.format_exc()}')
|
||||
|
||||
self._drop_pending_form_task(task_id, pending)
|
||||
|
||||
def _drop_pending_form_task(self, task_id: str, pending: dict) -> None:
|
||||
self._pending_forms_by_task.pop(task_id, None)
|
||||
msg_id = pending.get('msg_id', '')
|
||||
if msg_id:
|
||||
self._task_id_by_msg.pop(msg_id, None)
|
||||
self._stream_sessions.pop(msg_id, None)
|
||||
|
||||
async def _dispatch_event(self, event: wecombotevent.WecomBotEvent):
|
||||
"""Dispatch a message event to registered handlers with deduplication."""
|
||||
try:
|
||||
|
||||
@@ -138,39 +138,6 @@ class MonitoringRouterGroup(group.RouterGroup):
|
||||
}
|
||||
)
|
||||
|
||||
@self.route('/tool-calls', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def get_tool_calls() -> str:
|
||||
"""Get tool call records"""
|
||||
bot_ids = quart.request.args.getlist('botId')
|
||||
pipeline_ids = quart.request.args.getlist('pipelineId')
|
||||
session_ids = quart.request.args.getlist('sessionId')
|
||||
start_time_str = quart.request.args.get('startTime')
|
||||
end_time_str = quart.request.args.get('endTime')
|
||||
limit = int(quart.request.args.get('limit', 100))
|
||||
offset = int(quart.request.args.get('offset', 0))
|
||||
|
||||
start_time = parse_iso_datetime(start_time_str)
|
||||
end_time = parse_iso_datetime(end_time_str)
|
||||
|
||||
tool_calls, total = await self.ap.monitoring_service.get_tool_calls(
|
||||
bot_ids=bot_ids if bot_ids else None,
|
||||
pipeline_ids=pipeline_ids if pipeline_ids else None,
|
||||
session_ids=session_ids if session_ids else None,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
return self.success(
|
||||
data={
|
||||
'tool_calls': tool_calls,
|
||||
'total': total,
|
||||
'limit': limit,
|
||||
'offset': offset,
|
||||
}
|
||||
)
|
||||
|
||||
@self.route('/embedding-calls', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def get_embedding_calls() -> str:
|
||||
"""Get embedding call records"""
|
||||
@@ -317,16 +284,6 @@ class MonitoringRouterGroup(group.RouterGroup):
|
||||
offset=0,
|
||||
)
|
||||
|
||||
# Get tool calls
|
||||
tool_calls, tool_calls_total = await self.ap.monitoring_service.get_tool_calls(
|
||||
bot_ids=bot_ids if bot_ids else None,
|
||||
pipeline_ids=pipeline_ids if pipeline_ids else None,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
limit=limit,
|
||||
offset=0,
|
||||
)
|
||||
|
||||
# Get sessions
|
||||
sessions, sessions_total = await self.ap.monitoring_service.get_sessions(
|
||||
bot_ids=bot_ids if bot_ids else None,
|
||||
@@ -361,14 +318,12 @@ class MonitoringRouterGroup(group.RouterGroup):
|
||||
'overview': overview,
|
||||
'messages': messages,
|
||||
'llmCalls': llm_calls,
|
||||
'toolCalls': tool_calls,
|
||||
'embeddingCalls': embedding_calls,
|
||||
'sessions': sessions,
|
||||
'errors': errors,
|
||||
'totalCount': {
|
||||
'messages': messages_total,
|
||||
'llmCalls': llm_calls_total,
|
||||
'toolCalls': tool_calls_total,
|
||||
'embeddingCalls': embedding_calls_total,
|
||||
'sessions': sessions_total,
|
||||
'errors': errors_total,
|
||||
|
||||
@@ -62,16 +62,24 @@ class EmbedRouterGroup(group.RouterGroup):
|
||||
"""Resolve *bot_uuid* to ``(runtime_bot, pipeline_uuid)``.
|
||||
|
||||
Returns ``(None, None)`` when the bot does not exist, is not a
|
||||
``web_page_bot``, is disabled, or has no pipeline bound.
|
||||
``web_page_bot``, is disabled, or has no pipeline/workflow bound.
|
||||
"""
|
||||
for bot in self.ap.platform_mgr.bots:
|
||||
if (
|
||||
bot.bot_entity.uuid == bot_uuid
|
||||
and bot.bot_entity.adapter == 'web_page_bot'
|
||||
and bot.bot_entity.enable
|
||||
and bot.bot_entity.use_pipeline_uuid
|
||||
):
|
||||
return bot, bot.bot_entity.use_pipeline_uuid
|
||||
# Check for workflow binding first
|
||||
binding_type = getattr(bot.bot_entity, 'binding_type', 'pipeline') or 'pipeline'
|
||||
binding_uuid = getattr(bot.bot_entity, 'binding_uuid', None)
|
||||
|
||||
if binding_type == 'workflow' and binding_uuid:
|
||||
# For workflow binding, return workflow UUID
|
||||
return bot, binding_uuid
|
||||
elif bot.bot_entity.use_pipeline_uuid:
|
||||
# For pipeline binding, return pipeline UUID
|
||||
return bot, bot.bot_entity.use_pipeline_uuid
|
||||
return None, None
|
||||
|
||||
def _get_bot_config(self, bot_uuid: str) -> dict:
|
||||
|
||||
@@ -5,29 +5,6 @@ from ... import group
|
||||
from langbot.pkg.utils import importutil
|
||||
|
||||
|
||||
def _decrypt_qqofficial_secret(encrypted_b64: str, key: bytes) -> str:
|
||||
"""Decrypt the AppSecret returned by the QQ Official QR binding endpoint.
|
||||
|
||||
The base64 payload is laid out as `nonce (12 B) | ciphertext | tag (16 B)`.
|
||||
`key` is the 32-byte AES-256 key locally generated when the bind task
|
||||
was created and submitted as `key` to `q.qq.com/lite/create_bind_task`.
|
||||
"""
|
||||
import base64
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
try:
|
||||
raw = base64.b64decode(encrypted_b64)
|
||||
except Exception as exc:
|
||||
raise ValueError('Malformed encrypted credential') from exc
|
||||
if len(key) != 32 or len(raw) <= 28:
|
||||
raise ValueError('Invalid encrypted credential layout')
|
||||
nonce, ciphertext, tag = raw[:12], raw[12:-16], raw[-16:]
|
||||
try:
|
||||
return AESGCM(key).decrypt(nonce, ciphertext + tag, None).decode('utf-8')
|
||||
except Exception as exc:
|
||||
raise ValueError('Failed to decrypt credential') from exc
|
||||
|
||||
|
||||
@group.group_class('adapters', '/api/v1/platform/adapters')
|
||||
class AdaptersRouterGroup(group.RouterGroup):
|
||||
async def initialize(self) -> None:
|
||||
@@ -60,15 +37,6 @@ class AdaptersRouterGroup(group.RouterGroup):
|
||||
importutil.read_resource_file_bytes(icon_path), mimetype=mimetypes.guess_type(icon_path)[0]
|
||||
)
|
||||
|
||||
@self.route('/dingtalk/human-input-card-template', methods=['GET'], auth_type=group.AuthType.NONE)
|
||||
async def _() -> quart.Response:
|
||||
filename = 'dingtalk_human_input_card.json'
|
||||
response = quart.Response(
|
||||
importutil.read_resource_file_bytes(f'templates/{filename}'), mimetype='application/json'
|
||||
)
|
||||
response.headers['Content-Disposition'] = f'attachment; filename={filename}'
|
||||
return response
|
||||
|
||||
# In-memory session store for active registrations
|
||||
_create_app_sessions: dict = {}
|
||||
_SESSION_TTL = 900 # 15 minutes
|
||||
@@ -682,220 +650,3 @@ class AdaptersRouterGroup(group.RouterGroup):
|
||||
if session and session.get('task') and not session['task'].done():
|
||||
session['task'].cancel()
|
||||
return self.success(data={})
|
||||
|
||||
# -----------------------------------------------------------------------
|
||||
# QQ Official QR Binding
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
_qqofficial_sessions: dict = {}
|
||||
_QQOFFICIAL_SESSION_TTL = 300 # 5 minutes (QQ bind QR validity window)
|
||||
|
||||
def _cleanup_expired_qqofficial_sessions():
|
||||
import time
|
||||
|
||||
now = time.time()
|
||||
expired = [
|
||||
sid for sid, s in _qqofficial_sessions.items() if now - s.get('created_at', 0) > _QQOFFICIAL_SESSION_TTL
|
||||
]
|
||||
for sid in expired:
|
||||
session = _qqofficial_sessions.pop(sid, None)
|
||||
if session and session.get('task') and not session['task'].done():
|
||||
session['task'].cancel()
|
||||
|
||||
@self.route('/qqofficial/bind', methods=['POST'])
|
||||
async def _() -> str:
|
||||
"""Start QQ Official QR binding. Returns session_id + QR URL.
|
||||
|
||||
Flow: generate a local AES-256 key, register it with
|
||||
`q.qq.com/lite/create_bind_task`, then poll
|
||||
`q.qq.com/lite/poll_bind_result` until the user authorizes the
|
||||
bind inside the QQ Bot Assistant on mobile QQ. The encrypted
|
||||
AppSecret returned by the poll endpoint is decrypted with the
|
||||
same key. The key never leaves this process.
|
||||
"""
|
||||
import uuid
|
||||
import time
|
||||
import secrets
|
||||
import base64
|
||||
import aiohttp
|
||||
|
||||
QQ_BIND_BASE = 'https://q.qq.com'
|
||||
_cleanup_expired_qqofficial_sessions()
|
||||
|
||||
bind_key_bytes = secrets.token_bytes(32)
|
||||
bind_key = base64.b64encode(bind_key_bytes).decode('ascii')
|
||||
|
||||
session_id = str(uuid.uuid4())
|
||||
session = {
|
||||
'status': 'pending',
|
||||
'qr_url': None,
|
||||
'expire_at': None,
|
||||
'appid': None,
|
||||
'secret': None,
|
||||
'user_openid': None,
|
||||
'error': None,
|
||||
'created_at': time.time(),
|
||||
'task_id': None,
|
||||
'bind_key_bytes': bind_key_bytes,
|
||||
'interval': 2,
|
||||
}
|
||||
_qqofficial_sessions[session_id] = session
|
||||
|
||||
async def run_qr_binding():
|
||||
try:
|
||||
timeout = aiohttp.ClientTimeout(total=10)
|
||||
async with aiohttp.ClientSession(timeout=timeout) as http:
|
||||
# Step 1: create_bind_task — register our AES key, get task_id
|
||||
async with http.post(
|
||||
f'{QQ_BIND_BASE}/lite/create_bind_task',
|
||||
json={'key': bind_key},
|
||||
headers={'Accept': 'application/json'},
|
||||
) as resp:
|
||||
try:
|
||||
data = await resp.json(content_type=None)
|
||||
except (aiohttp.ContentTypeError, ValueError):
|
||||
session['status'] = 'error'
|
||||
session['error'] = 'Invalid response from QQ bind service'
|
||||
return
|
||||
if int(data.get('retcode', -1)) != 0:
|
||||
session['status'] = 'error'
|
||||
session['error'] = (
|
||||
data.get('msg') or data.get('message') or 'Failed to create bind task'
|
||||
)
|
||||
return
|
||||
task_id = str((data.get('data') or {}).get('task_id') or '').strip()
|
||||
if not task_id:
|
||||
session['status'] = 'error'
|
||||
session['error'] = 'Missing task_id in QQ response'
|
||||
return
|
||||
|
||||
# The QR encodes a URL that mobile QQ opens inside the QQ Bot Assistant.
|
||||
# `source=langbot` is a courtesy attribution parameter so Tencent
|
||||
# can see LangBot adoption metrics, matching the convention used by
|
||||
# other third-party integrations (e.g. hermes-agent uses `source=hermes`).
|
||||
qr_url = f'{QQ_BIND_BASE}/qqbot/openclaw/connect.html?task_id={task_id}&_wv=2&source=langbot'
|
||||
session['task_id'] = task_id
|
||||
session['qr_url'] = qr_url
|
||||
session['expire_at'] = time.time() + _QQOFFICIAL_SESSION_TTL
|
||||
session['status'] = 'waiting'
|
||||
|
||||
# Step 2: poll_bind_result until completed (status=2) or expired (3).
|
||||
deadline = time.time() + _QQOFFICIAL_SESSION_TTL
|
||||
while time.time() < deadline:
|
||||
await asyncio.sleep(session['interval'])
|
||||
|
||||
async with http.post(
|
||||
f'{QQ_BIND_BASE}/lite/poll_bind_result',
|
||||
json={'task_id': task_id},
|
||||
headers={'Accept': 'application/json'},
|
||||
) as poll_resp:
|
||||
try:
|
||||
poll_data = await poll_resp.json(content_type=None)
|
||||
except (aiohttp.ContentTypeError, ValueError):
|
||||
continue
|
||||
|
||||
if int(poll_data.get('retcode', -1)) != 0:
|
||||
session['status'] = 'error'
|
||||
session['error'] = poll_data.get('msg') or poll_data.get('message') or 'Poll failed'
|
||||
return
|
||||
|
||||
payload = poll_data.get('data') or {}
|
||||
try:
|
||||
raw_status = int(payload.get('status', 0))
|
||||
except (TypeError, ValueError):
|
||||
raw_status = 0
|
||||
|
||||
if raw_status == 2:
|
||||
appid = str(payload.get('bot_appid') or '').strip()
|
||||
encrypted = str(payload.get('bot_encrypt_secret') or '').strip()
|
||||
if not appid or not encrypted:
|
||||
session['status'] = 'error'
|
||||
session['error'] = 'Incomplete credential payload'
|
||||
return
|
||||
try:
|
||||
session['secret'] = _decrypt_qqofficial_secret(
|
||||
encrypted,
|
||||
bind_key_bytes,
|
||||
)
|
||||
except ValueError as exc:
|
||||
session['status'] = 'error'
|
||||
session['error'] = str(exc)
|
||||
return
|
||||
session['appid'] = appid
|
||||
# The scanner's OpenID is returned alongside the credentials —
|
||||
# surfaced to the dashboard for audit / "bound by" display.
|
||||
session['user_openid'] = str(payload.get('user_openid') or '').strip() or None
|
||||
session['status'] = 'success'
|
||||
return
|
||||
|
||||
if raw_status == 3:
|
||||
session['status'] = 'expired'
|
||||
session['error'] = 'QR code expired'
|
||||
return
|
||||
# status 0 / 1: still pending, continue polling
|
||||
|
||||
session['status'] = 'expired'
|
||||
session['error'] = 'QR code expired'
|
||||
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
except Exception as e:
|
||||
session['status'] = 'error'
|
||||
session['error'] = str(e)
|
||||
|
||||
task = asyncio.create_task(run_qr_binding())
|
||||
session['task'] = task
|
||||
|
||||
# Wait up to 10s for the QR URL to be ready before responding.
|
||||
for _ in range(20):
|
||||
if session['qr_url'] or session['error']:
|
||||
break
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
if session['error']:
|
||||
task.cancel()
|
||||
return self.http_status(502, -1, session['error'])
|
||||
|
||||
if not session['qr_url']:
|
||||
task.cancel()
|
||||
session['status'] = 'error'
|
||||
session['error'] = 'Timeout waiting for QR code'
|
||||
return self.http_status(504, -1, 'Timeout waiting for QR code')
|
||||
|
||||
return self.success(
|
||||
data={
|
||||
'session_id': session_id,
|
||||
'qr_url': session['qr_url'],
|
||||
'expire_at': session['expire_at'],
|
||||
}
|
||||
)
|
||||
|
||||
@self.route('/qqofficial/bind/status/<session_id>', methods=['GET'])
|
||||
async def _(session_id: str) -> str:
|
||||
"""Poll QQ Official QR binding status."""
|
||||
_cleanup_expired_qqofficial_sessions()
|
||||
session = _qqofficial_sessions.get(session_id)
|
||||
if not session:
|
||||
return self.http_status(404, -1, 'Session not found')
|
||||
|
||||
data = {'status': session['status']}
|
||||
|
||||
if session['status'] == 'success':
|
||||
data['appid'] = session['appid']
|
||||
data['secret'] = session['secret']
|
||||
if session.get('user_openid'):
|
||||
data['user_openid'] = session['user_openid']
|
||||
_qqofficial_sessions.pop(session_id, None)
|
||||
elif session['status'] in ('error', 'expired'):
|
||||
data['error'] = session['error']
|
||||
_qqofficial_sessions.pop(session_id, None)
|
||||
|
||||
return self.success(data=data)
|
||||
|
||||
@self.route('/qqofficial/bind/<session_id>', methods=['DELETE'])
|
||||
async def _(session_id: str) -> str:
|
||||
"""Cancel and clean up a QQ Official QR binding session."""
|
||||
session = _qqofficial_sessions.pop(session_id, None)
|
||||
if session and session.get('task') and not session['task'].done():
|
||||
session['task'].cancel()
|
||||
return self.success(data={})
|
||||
|
||||
@@ -29,11 +29,11 @@ class MCPRouterGroup(group.RouterGroup):
|
||||
traceback.print_exc()
|
||||
return self.http_status(500, -1, f'Failed to create MCP server: {str(e)}')
|
||||
|
||||
@self.route(
|
||||
'/servers/<path:server_name>', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN
|
||||
)
|
||||
@self.route('/servers/<server_name>', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def _(server_name: str) -> str:
|
||||
"""获取、更新或删除MCP服务器配置"""
|
||||
from urllib.parse import unquote
|
||||
|
||||
server_name = unquote(server_name)
|
||||
|
||||
server_data = await self.ap.mcp_service.get_mcp_server_by_name(server_name)
|
||||
@@ -58,15 +58,17 @@ class MCPRouterGroup(group.RouterGroup):
|
||||
except Exception as e:
|
||||
return self.http_status(500, -1, f'Failed to delete MCP server: {str(e)}')
|
||||
|
||||
@self.route('/servers/<path:server_name>/test', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
|
||||
@self.route('/servers/<server_name>/test', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def _(server_name: str) -> str:
|
||||
"""测试MCP服务器连接"""
|
||||
from urllib.parse import unquote
|
||||
|
||||
server_name = unquote(server_name)
|
||||
server_data = await quart.request.json
|
||||
task_id = await self.ap.mcp_service.test_mcp_server(server_name=server_name, server_data=server_data)
|
||||
return self.success(data={'task_id': task_id})
|
||||
|
||||
@self.route('/servers/<path:server_name>/resources', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
|
||||
@self.route('/servers/<server_name>/resources', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def _(server_name: str) -> str:
|
||||
"""Get resources from an MCP server"""
|
||||
server_name = unquote(server_name)
|
||||
@@ -84,9 +86,7 @@ class MCPRouterGroup(group.RouterGroup):
|
||||
except Exception as e:
|
||||
return self.http_status(500, -1, f'Failed to get resources: {str(e)}')
|
||||
|
||||
@self.route(
|
||||
'/servers/<path:server_name>/resource-templates', methods=['GET'], auth_type=group.AuthType.USER_TOKEN
|
||||
)
|
||||
@self.route('/servers/<server_name>/resource-templates', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def _(server_name: str) -> str:
|
||||
"""Get resource templates from an MCP server"""
|
||||
server_name = unquote(server_name)
|
||||
@@ -96,20 +96,7 @@ class MCPRouterGroup(group.RouterGroup):
|
||||
except Exception as e:
|
||||
return self.http_status(500, -1, f'Failed to get resource templates: {str(e)}')
|
||||
|
||||
@self.route('/servers/<path:server_name>/logs', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def _(server_name: str) -> str:
|
||||
"""Get logs from an MCP server"""
|
||||
server_name = unquote(server_name)
|
||||
try:
|
||||
limit = int(quart.request.args.get('limit', 200))
|
||||
except (TypeError, ValueError):
|
||||
limit = 200
|
||||
limit = min(limit, 500)
|
||||
level = quart.request.args.get('level') or None
|
||||
logs = await self.ap.mcp_service.get_mcp_server_logs(server_name, limit=limit, level=level)
|
||||
return self.success(data={'logs': logs})
|
||||
|
||||
@self.route('/servers/<path:server_name>/resources/read', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
|
||||
@self.route('/servers/<server_name>/resources/read', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def _(server_name: str) -> str:
|
||||
"""Read a resource from an MCP server"""
|
||||
server_name = unquote(server_name)
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
# Workflow router group
|
||||
from .workflows import WorkflowsRouterGroup, ExecutionsRouterGroup
|
||||
from .websocket_chat import WorkflowWebSocketChatRouterGroup
|
||||
|
||||
__all__ = ['WorkflowsRouterGroup', 'ExecutionsRouterGroup', 'WorkflowWebSocketChatRouterGroup']
|
||||
@@ -0,0 +1,260 @@
|
||||
"""Workflow WebSocket聊天路由 - 支持工作流调试的双向实时通信"""
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
import json
|
||||
import logging
|
||||
|
||||
import quart
|
||||
|
||||
from ... import group
|
||||
from ......platform.sources.websocket_manager import ws_connection_manager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@group.group_class('workflow_websocket_chat', '/api/v1/workflows/<workflow_uuid>/ws')
|
||||
class WorkflowWebSocketChatRouterGroup(group.RouterGroup):
|
||||
async def initialize(self) -> None:
|
||||
@self.quart_app.websocket(self.path + '/connect')
|
||||
async def workflow_websocket_connect(workflow_uuid: str):
|
||||
"""
|
||||
建立工作流WebSocket连接
|
||||
|
||||
URL参数:
|
||||
- workflow_uuid: 工作流UUID
|
||||
- session_type: 会话类型 (person/group)
|
||||
"""
|
||||
try:
|
||||
session_type = quart.websocket.args.get('session_type', 'person')
|
||||
logger.info(
|
||||
'Workflow WebSocket connect request received',
|
||||
extra={
|
||||
'workflow_uuid': workflow_uuid,
|
||||
'session_type': session_type,
|
||||
'path': quart.websocket.path,
|
||||
'query_string': quart.websocket.query_string.decode('utf-8', errors='ignore'),
|
||||
'remote_addr': getattr(quart.websocket, 'remote_addr', None),
|
||||
'user_agent': quart.websocket.headers.get('User-Agent', ''),
|
||||
'host': quart.websocket.headers.get('Host', ''),
|
||||
'origin': quart.websocket.headers.get('Origin', ''),
|
||||
},
|
||||
)
|
||||
|
||||
if session_type not in ['person', 'group']:
|
||||
await quart.websocket.send(
|
||||
json.dumps({'type': 'error', 'message': 'session_type must be person or group'})
|
||||
)
|
||||
return
|
||||
|
||||
websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
|
||||
|
||||
if not websocket_adapter:
|
||||
logger.warning(
|
||||
'Workflow WebSocket adapter missing',
|
||||
extra={
|
||||
'workflow_uuid': workflow_uuid,
|
||||
'session_type': session_type,
|
||||
},
|
||||
)
|
||||
await quart.websocket.send(json.dumps({'type': 'error', 'message': 'WebSocket adapter not found'}))
|
||||
return
|
||||
|
||||
connection = await ws_connection_manager.add_connection(
|
||||
websocket=quart.websocket._get_current_object(),
|
||||
pipeline_uuid=workflow_uuid,
|
||||
session_type=session_type,
|
||||
metadata={'user_agent': quart.websocket.headers.get('User-Agent', ''), 'is_workflow': True},
|
||||
)
|
||||
|
||||
await quart.websocket.send(
|
||||
json.dumps(
|
||||
{
|
||||
'type': 'connected',
|
||||
'connection_id': connection.connection_id,
|
||||
'workflow_uuid': workflow_uuid,
|
||||
'session_type': session_type,
|
||||
'timestamp': connection.created_at.isoformat(),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f'Workflow WebSocket connection established: {connection.connection_id} '
|
||||
f'(workflow={workflow_uuid}, session_type={session_type})'
|
||||
)
|
||||
|
||||
receive_task = asyncio.create_task(self._handle_receive(connection, websocket_adapter))
|
||||
send_task = asyncio.create_task(self._handle_send(connection))
|
||||
|
||||
try:
|
||||
await asyncio.gather(receive_task, send_task)
|
||||
except Exception as e:
|
||||
logger.error(f'Workflow WebSocket task execution error: {e}')
|
||||
finally:
|
||||
await ws_connection_manager.remove_connection(connection.connection_id)
|
||||
logger.debug(f'Workflow WebSocket connection cleaned: {connection.connection_id}')
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
'Workflow WebSocket connection error',
|
||||
exc_info=True,
|
||||
extra={
|
||||
'workflow_uuid': workflow_uuid,
|
||||
'session_type': quart.websocket.args.get('session_type', 'person'),
|
||||
'path': quart.websocket.path,
|
||||
'query_string': quart.websocket.query_string.decode('utf-8', errors='ignore'),
|
||||
'remote_addr': getattr(quart.websocket, 'remote_addr', None),
|
||||
},
|
||||
)
|
||||
try:
|
||||
await quart.websocket.send(json.dumps({'type': 'error', 'message': str(e)}))
|
||||
except Exception as send_error:
|
||||
logger.debug(
|
||||
'Failed to send error message to workflow websocket client',
|
||||
exc_info=True,
|
||||
extra={
|
||||
'workflow_uuid': workflow_uuid,
|
||||
'send_error': str(send_error),
|
||||
},
|
||||
)
|
||||
|
||||
@self.route('/messages/<session_type>', methods=['GET'])
|
||||
async def get_messages(workflow_uuid: str, session_type: str) -> str:
|
||||
"""获取工作流消息历史"""
|
||||
try:
|
||||
if session_type not in ['person', 'group']:
|
||||
return self.http_status(400, -1, 'session_type must be person or group')
|
||||
|
||||
websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
|
||||
|
||||
if not websocket_adapter:
|
||||
return self.http_status(404, -1, 'WebSocket adapter not found')
|
||||
|
||||
messages = websocket_adapter.get_websocket_messages(workflow_uuid, session_type)
|
||||
|
||||
return self.success(data={'messages': messages})
|
||||
|
||||
except Exception as e:
|
||||
return self.http_status(500, -1, f'Internal server error: {str(e)}')
|
||||
|
||||
@self.route('/reset/<session_type>', methods=['POST'])
|
||||
async def reset_session(workflow_uuid: str, session_type: str) -> str:
|
||||
"""重置工作流会话"""
|
||||
try:
|
||||
if session_type not in ['person', 'group']:
|
||||
return self.http_status(400, -1, 'session_type must be person or group')
|
||||
|
||||
websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
|
||||
|
||||
if not websocket_adapter:
|
||||
return self.http_status(404, -1, 'WebSocket adapter not found')
|
||||
|
||||
websocket_adapter.reset_session(workflow_uuid, session_type)
|
||||
|
||||
return self.success(data={'message': 'Session reset successfully'})
|
||||
|
||||
except Exception as e:
|
||||
return self.http_status(500, -1, f'Internal server error: {str(e)}')
|
||||
|
||||
@self.route('/connections', methods=['GET'])
|
||||
async def get_connections(workflow_uuid: str) -> str:
|
||||
"""获取当前工作流连接统计"""
|
||||
try:
|
||||
stats = ws_connection_manager.get_stats()
|
||||
connections = await ws_connection_manager.get_connections_by_pipeline(workflow_uuid)
|
||||
|
||||
return self.success(
|
||||
data={
|
||||
'stats': stats,
|
||||
'connections': [
|
||||
{
|
||||
'connection_id': conn.connection_id,
|
||||
'session_type': conn.session_type,
|
||||
'created_at': conn.created_at.isoformat(),
|
||||
'last_active': conn.last_active.isoformat(),
|
||||
'is_active': conn.is_active,
|
||||
}
|
||||
for conn in connections
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
return self.http_status(500, -1, f'Internal server error: {str(e)}')
|
||||
|
||||
@self.route('/broadcast', methods=['POST'])
|
||||
async def broadcast_message(workflow_uuid: str) -> str:
|
||||
"""向所有工作流连接广播消息"""
|
||||
try:
|
||||
data = await quart.request.get_json()
|
||||
message = data.get('message')
|
||||
|
||||
if not message:
|
||||
return self.http_status(400, -1, 'message is required')
|
||||
|
||||
broadcast_data = {
|
||||
'type': 'broadcast',
|
||||
'message': message,
|
||||
'timestamp': datetime.datetime.now().isoformat(),
|
||||
}
|
||||
|
||||
await ws_connection_manager.broadcast_to_pipeline(workflow_uuid, broadcast_data)
|
||||
|
||||
return self.success(data={'message': 'Broadcast sent successfully'})
|
||||
|
||||
except Exception as e:
|
||||
return self.http_status(500, -1, f'Internal server error: {str(e)}')
|
||||
|
||||
async def _handle_receive(self, connection, websocket_adapter):
|
||||
"""处理接收消息的任务"""
|
||||
try:
|
||||
while connection.is_active:
|
||||
message = await quart.websocket.receive()
|
||||
|
||||
await ws_connection_manager.update_activity(connection.connection_id)
|
||||
|
||||
try:
|
||||
data = json.loads(message)
|
||||
message_type = data.get('type', 'message')
|
||||
|
||||
if message_type == 'ping':
|
||||
await connection.send_queue.put(
|
||||
{'type': 'pong', 'timestamp': datetime.datetime.now().isoformat()}
|
||||
)
|
||||
|
||||
elif message_type == 'message':
|
||||
logger.debug(f'收到工作流消息: {data} from {connection.connection_id}')
|
||||
await websocket_adapter.handle_websocket_message(connection, data)
|
||||
|
||||
elif message_type == 'disconnect':
|
||||
logger.debug(f'Client disconnected: {connection.connection_id}')
|
||||
break
|
||||
|
||||
else:
|
||||
logger.warning(f'Unknown message type: {message_type}')
|
||||
|
||||
except json.JSONDecodeError:
|
||||
logger.error(f'Invalid JSON message: {message}')
|
||||
await connection.send_queue.put({'type': 'error', 'message': 'Invalid JSON format'})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f'Receive message error: {e}', exc_info=True)
|
||||
finally:
|
||||
connection.is_active = False
|
||||
|
||||
async def _handle_send(self, connection):
|
||||
"""处理发送消息的任务"""
|
||||
try:
|
||||
while connection.is_active:
|
||||
try:
|
||||
message = await asyncio.wait_for(connection.send_queue.get(), timeout=1.0)
|
||||
await quart.websocket.send(json.dumps(message))
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f'Send message error: {e}', exc_info=True)
|
||||
finally:
|
||||
connection.is_active = False
|
||||
@@ -0,0 +1,484 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import quart
|
||||
|
||||
from ... import group
|
||||
from ....service.workflow import WorkflowExecutionFailedError
|
||||
|
||||
|
||||
@group.group_class('workflows', '/api/v1/workflows')
|
||||
class WorkflowsRouterGroup(group.RouterGroup):
|
||||
"""Workflow API router group"""
|
||||
|
||||
async def initialize(self) -> None:
|
||||
# Workflow CRUD
|
||||
@self.route('', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _() -> str:
|
||||
if quart.request.method == 'GET':
|
||||
sort_by = quart.request.args.get('sort_by', 'created_at')
|
||||
sort_order = quart.request.args.get('sort_order', 'DESC')
|
||||
enabled_only = quart.request.args.get('enabled_only', 'false').lower() == 'true'
|
||||
return self.success(
|
||||
data={'workflows': await self.ap.workflow_service.get_workflows(sort_by, sort_order, enabled_only)}
|
||||
)
|
||||
elif quart.request.method == 'POST':
|
||||
json_data = await quart.request.json
|
||||
workflow_uuid = await self.ap.workflow_service.create_workflow(json_data)
|
||||
return self.success(data={'uuid': workflow_uuid})
|
||||
|
||||
# Get node types (available nodes for the editor)
|
||||
@self.route('/_/node-types', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _() -> str:
|
||||
return self.success(
|
||||
data={
|
||||
'node_types': await self.ap.workflow_service.get_node_types(),
|
||||
'categories': await self.ap.workflow_service.get_node_types_by_category_meta(),
|
||||
}
|
||||
)
|
||||
|
||||
# Get node types by category
|
||||
@self.route('/_/node-types/categories', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _() -> str:
|
||||
return self.success(data={'categories': await self.ap.workflow_service.get_node_types_by_category()})
|
||||
|
||||
# Single workflow operations
|
||||
@self.route(
|
||||
'/<workflow_uuid>', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
|
||||
)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
if quart.request.method == 'GET':
|
||||
workflow = await self.ap.workflow_service.get_workflow(workflow_uuid)
|
||||
if workflow is None:
|
||||
return self.http_status(404, -1, 'workflow not found')
|
||||
return self.success(data={'workflow': workflow})
|
||||
elif quart.request.method == 'PUT':
|
||||
json_data = await quart.request.json
|
||||
try:
|
||||
await self.ap.workflow_service.update_workflow(workflow_uuid, json_data)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
elif quart.request.method == 'DELETE':
|
||||
await self.ap.workflow_service.delete_workflow(workflow_uuid)
|
||||
return self.success()
|
||||
return self.http_status(405, -1, 'method not allowed')
|
||||
|
||||
# Publish workflow (enable)
|
||||
@self.route('/<workflow_uuid>/publish', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
try:
|
||||
await self.ap.workflow_service.publish_workflow(workflow_uuid)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Unpublish workflow (disable)
|
||||
@self.route('/<workflow_uuid>/unpublish', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
try:
|
||||
await self.ap.workflow_service.unpublish_workflow(workflow_uuid)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Copy workflow
|
||||
@self.route('/<workflow_uuid>/copy', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
try:
|
||||
new_uuid = await self.ap.workflow_service.copy_workflow(workflow_uuid)
|
||||
return self.success(data={'uuid': new_uuid})
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Execute workflow manually
|
||||
@self.route('/<workflow_uuid>/execute', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
json_data = await quart.request.json or {}
|
||||
trigger_data = json_data.get('trigger_data', {})
|
||||
session_id = json_data.get('session_id')
|
||||
user_id = json_data.get('user_id')
|
||||
bot_id = json_data.get('bot_id')
|
||||
|
||||
try:
|
||||
execution_id = await self.ap.workflow_service.execute_workflow(
|
||||
workflow_uuid,
|
||||
trigger_type='manual',
|
||||
trigger_data=trigger_data,
|
||||
session_id=session_id,
|
||||
user_id=user_id,
|
||||
bot_id=bot_id,
|
||||
)
|
||||
return self.success(data={'execution_id': execution_id})
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
except WorkflowExecutionFailedError as e:
|
||||
return self.http_status(500, -1, e.message)
|
||||
|
||||
# Get workflow executions
|
||||
@self.route('/<workflow_uuid>/executions', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
limit = int(quart.request.args.get('limit', 50))
|
||||
offset = int(quart.request.args.get('offset', 0))
|
||||
executions = await self.ap.workflow_service.get_executions(
|
||||
workflow_uuid=workflow_uuid, limit=limit, offset=offset
|
||||
)
|
||||
return self.success(data=executions)
|
||||
|
||||
@self.route(
|
||||
'/<workflow_uuid>/executions/<execution_uuid>',
|
||||
methods=['GET'],
|
||||
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
|
||||
)
|
||||
async def _(workflow_uuid: str, execution_uuid: str) -> str:
|
||||
execution = await self.ap.workflow_service.get_execution(execution_uuid)
|
||||
if execution is None:
|
||||
return self.http_status(404, -1, 'execution not found')
|
||||
if execution.get('workflow_uuid') != workflow_uuid:
|
||||
return self.http_status(404, -1, 'execution not found in workflow')
|
||||
return self.success(data={'execution': execution})
|
||||
|
||||
# Get workflow versions
|
||||
@self.route('/<workflow_uuid>/versions', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
versions = await self.ap.workflow_service.get_versions(workflow_uuid)
|
||||
return self.success(data={'versions': versions})
|
||||
|
||||
# Rollback to a specific version
|
||||
@self.route(
|
||||
'/<workflow_uuid>/rollback/<int:version>', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
|
||||
)
|
||||
async def _(workflow_uuid: str, version: int) -> str:
|
||||
try:
|
||||
await self.ap.workflow_service.rollback_to_version(workflow_uuid, version)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Workflow extensions (plugins and MCP servers)
|
||||
@self.route(
|
||||
'/<workflow_uuid>/extensions', methods=['GET', 'PUT'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
|
||||
)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
if quart.request.method == 'GET':
|
||||
workflow = await self.ap.workflow_service.get_workflow(workflow_uuid)
|
||||
if workflow is None:
|
||||
return self.http_status(404, -1, 'workflow not found')
|
||||
|
||||
# Get available plugins and MCP servers
|
||||
pipeline_component_kinds = ['Command', 'EventListener', 'Tool']
|
||||
plugins = await self.ap.plugin_connector.list_plugins(component_kinds=pipeline_component_kinds)
|
||||
mcp_servers = await self.ap.mcp_service.get_mcp_servers(contain_runtime_info=True)
|
||||
|
||||
extensions_prefs = workflow.get('extensions_preferences', {})
|
||||
return self.success(
|
||||
data={
|
||||
'enable_all_plugins': extensions_prefs.get('enable_all_plugins', True),
|
||||
'enable_all_mcp_servers': extensions_prefs.get('enable_all_mcp_servers', True),
|
||||
'bound_plugins': extensions_prefs.get('plugins', []),
|
||||
'available_plugins': plugins,
|
||||
'bound_mcp_servers': extensions_prefs.get('mcp_servers', []),
|
||||
'available_mcp_servers': mcp_servers,
|
||||
}
|
||||
)
|
||||
elif quart.request.method == 'PUT':
|
||||
json_data = await quart.request.json
|
||||
enable_all_plugins = json_data.get('enable_all_plugins', True)
|
||||
enable_all_mcp_servers = json_data.get('enable_all_mcp_servers', True)
|
||||
bound_plugins = json_data.get('bound_plugins', [])
|
||||
bound_mcp_servers = json_data.get('bound_mcp_servers', [])
|
||||
|
||||
try:
|
||||
await self.ap.workflow_service.update_workflow_extensions(
|
||||
workflow_uuid, bound_plugins, bound_mcp_servers, enable_all_plugins, enable_all_mcp_servers
|
||||
)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
return self.http_status(405, -1, 'method not allowed')
|
||||
|
||||
# Debug API - Start debug execution
|
||||
@self.route('/<workflow_uuid>/debug/start', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
json_data = await quart.request.json or {}
|
||||
context = json_data.get('context', {})
|
||||
variables = json_data.get('variables', {})
|
||||
breakpoints = json_data.get('breakpoints', [])
|
||||
|
||||
try:
|
||||
execution_id = await self.ap.workflow_service.start_debug_execution(
|
||||
workflow_uuid, context=context, variables=variables, breakpoints=breakpoints
|
||||
)
|
||||
return self.success(data={'execution_id': execution_id})
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Debug API - Pause execution
|
||||
@self.route(
|
||||
'/<workflow_uuid>/debug/<execution_uuid>/pause',
|
||||
methods=['POST'],
|
||||
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
|
||||
)
|
||||
async def _(workflow_uuid: str, execution_uuid: str) -> str:
|
||||
try:
|
||||
await self.ap.workflow_service.pause_debug_execution(workflow_uuid, execution_uuid)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Debug API - Resume execution
|
||||
@self.route(
|
||||
'/<workflow_uuid>/debug/<execution_uuid>/resume',
|
||||
methods=['POST'],
|
||||
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
|
||||
)
|
||||
async def _(workflow_uuid: str, execution_uuid: str) -> str:
|
||||
try:
|
||||
await self.ap.workflow_service.resume_debug_execution(workflow_uuid, execution_uuid)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Debug API - Step execution
|
||||
@self.route(
|
||||
'/<workflow_uuid>/debug/<execution_uuid>/step',
|
||||
methods=['POST'],
|
||||
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
|
||||
)
|
||||
async def _(workflow_uuid: str, execution_uuid: str) -> str:
|
||||
try:
|
||||
result = await self.ap.workflow_service.step_debug_execution(workflow_uuid, execution_uuid)
|
||||
return self.success(data=result)
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Debug API - Stop execution
|
||||
@self.route(
|
||||
'/<workflow_uuid>/debug/<execution_uuid>/stop',
|
||||
methods=['POST'],
|
||||
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
|
||||
)
|
||||
async def _(workflow_uuid: str, execution_uuid: str) -> str:
|
||||
try:
|
||||
await self.ap.workflow_service.stop_debug_execution(workflow_uuid, execution_uuid)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Debug API - Get debug state
|
||||
@self.route(
|
||||
'/<workflow_uuid>/debug/<execution_uuid>/state',
|
||||
methods=['GET'],
|
||||
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
|
||||
)
|
||||
async def _(workflow_uuid: str, execution_uuid: str) -> str:
|
||||
try:
|
||||
state = await self.ap.workflow_service.get_debug_state(workflow_uuid, execution_uuid)
|
||||
return self.success(data=state)
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Get execution logs
|
||||
@self.route(
|
||||
'/<workflow_uuid>/executions/<execution_uuid>/logs',
|
||||
methods=['GET'],
|
||||
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
|
||||
)
|
||||
async def _(workflow_uuid: str, execution_uuid: str) -> str:
|
||||
limit = int(quart.request.args.get('limit', 100))
|
||||
offset = int(quart.request.args.get('offset', 0))
|
||||
try:
|
||||
result = await self.ap.workflow_service.get_execution_logs(workflow_uuid, execution_uuid, limit, offset)
|
||||
return self.success(data=result)
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Rerun execution
|
||||
@self.route(
|
||||
'/<workflow_uuid>/executions/<execution_uuid>/rerun',
|
||||
methods=['POST'],
|
||||
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
|
||||
)
|
||||
async def _(workflow_uuid: str, execution_uuid: str) -> str:
|
||||
try:
|
||||
new_execution_id = await self.ap.workflow_service.rerun_execution(workflow_uuid, execution_uuid)
|
||||
return self.success(data={'execution_uuid': new_execution_id})
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# Get workflow statistics
|
||||
@self.route('/<workflow_uuid>/stats', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(workflow_uuid: str) -> str:
|
||||
try:
|
||||
stats = await self.ap.workflow_service.get_workflow_stats(workflow_uuid)
|
||||
return self.success(data=stats)
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
|
||||
# LLM Node Performance Test Endpoint
|
||||
# Tests each step of LLM node execution with detailed timing
|
||||
@self.route('/_/test/llm-node', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _() -> str:
|
||||
"""Test LLM node performance with detailed step-by-step timing.
|
||||
|
||||
Request body:
|
||||
{
|
||||
"model_uuid": "uuid-of-model",
|
||||
"system_prompt": "optional system prompt",
|
||||
"user_prompt": "test message",
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 100
|
||||
}
|
||||
|
||||
Response includes timing for each step:
|
||||
- model_fetch: Time to get model from model_mgr
|
||||
- prompt_build: Time to build messages
|
||||
- llm_call: Time for actual LLM invocation
|
||||
- total: Total time
|
||||
- usage: Token usage information
|
||||
"""
|
||||
import time
|
||||
|
||||
json_data = await quart.request.json
|
||||
if not json_data:
|
||||
return self.http_status(400, -1, 'Request body is required')
|
||||
|
||||
model_uuid = json_data.get('model_uuid', '')
|
||||
if not model_uuid:
|
||||
return self.http_status(400, -1, 'model_uuid is required')
|
||||
|
||||
user_prompt = json_data.get('user_prompt', 'test')
|
||||
system_prompt = json_data.get('system_prompt', '')
|
||||
temperature = json_data.get('temperature')
|
||||
max_tokens = json_data.get('max_tokens', 0)
|
||||
|
||||
timings = {}
|
||||
errors = []
|
||||
|
||||
# Step 1: Model fetch
|
||||
t_start = time.perf_counter()
|
||||
try:
|
||||
runtime_model = await self.ap.model_mgr.get_model_by_uuid(model_uuid)
|
||||
timings['model_fetch_ms'] = round((time.perf_counter() - t_start) * 1000, 2)
|
||||
timings['model_found'] = True
|
||||
timings['model_name'] = runtime_model.model_entity.name if runtime_model else None
|
||||
except Exception as e:
|
||||
timings['model_fetch_ms'] = round((time.perf_counter() - t_start) * 1000, 2)
|
||||
timings['model_found'] = False
|
||||
errors.append(f'Model fetch failed: {str(e)}')
|
||||
return self.http_status(400, -1, {
|
||||
'error': errors[0],
|
||||
'timings': timings,
|
||||
})
|
||||
|
||||
# Step 2: Build messages
|
||||
t_start = time.perf_counter()
|
||||
import langbot_plugin.api.entities.builtin.provider.message as provider_message
|
||||
messages = []
|
||||
if system_prompt:
|
||||
messages.append(provider_message.Message(role='system', content=system_prompt))
|
||||
messages.append(provider_message.Message(role='user', content=user_prompt))
|
||||
timings['prompt_build_ms'] = round((time.perf_counter() - t_start) * 1000, 2)
|
||||
|
||||
# Step 3: Build extra args
|
||||
extra_args = {}
|
||||
if temperature is not None:
|
||||
extra_args['temperature'] = float(temperature)
|
||||
if max_tokens and int(max_tokens) > 0:
|
||||
extra_args['max_tokens'] = int(max_tokens)
|
||||
|
||||
# Step 4: LLM call
|
||||
t_start = time.perf_counter()
|
||||
try:
|
||||
result_message = await runtime_model.provider.invoke_llm(
|
||||
query=None,
|
||||
model=runtime_model,
|
||||
messages=messages,
|
||||
funcs=None,
|
||||
extra_args=extra_args,
|
||||
)
|
||||
timings['llm_call_ms'] = round((time.perf_counter() - t_start) * 1000, 2)
|
||||
timings['llm_call_success'] = True
|
||||
|
||||
# Extract response text
|
||||
response_text = ''
|
||||
if isinstance(result_message.content, str):
|
||||
response_text = result_message.content
|
||||
elif isinstance(result_message.content, list):
|
||||
for elem in result_message.content:
|
||||
if hasattr(elem, 'text') and elem.text:
|
||||
response_text += elem.text
|
||||
elif isinstance(elem, str):
|
||||
response_text += elem
|
||||
|
||||
timings['response_length'] = len(response_text)
|
||||
timings['response_preview'] = response_text[:200]
|
||||
|
||||
# Extract usage
|
||||
usage = {'prompt_tokens': 0, 'completion_tokens': 0, 'total_tokens': 0}
|
||||
if hasattr(result_message, 'usage') and result_message.usage:
|
||||
u = result_message.usage
|
||||
usage = {
|
||||
'prompt_tokens': getattr(u, 'prompt_tokens', 0) or 0,
|
||||
'completion_tokens': getattr(u, 'completion_tokens', 0) or 0,
|
||||
'total_tokens': getattr(u, 'total_tokens', 0) or 0,
|
||||
}
|
||||
timings['usage'] = usage
|
||||
|
||||
except Exception as e:
|
||||
timings['llm_call_ms'] = round((time.perf_counter() - t_start) * 1000, 2)
|
||||
timings['llm_call_success'] = False
|
||||
errors.append(f'LLM call failed: {str(e)}')
|
||||
|
||||
# Calculate total
|
||||
timings['total_ms'] = round(sum([
|
||||
timings.get('model_fetch_ms', 0),
|
||||
timings.get('prompt_build_ms', 0),
|
||||
timings.get('llm_call_ms', 0),
|
||||
]), 2)
|
||||
|
||||
# Add breakdown percentage
|
||||
if timings['total_ms'] > 0:
|
||||
timings['breakdown'] = {
|
||||
'model_fetch_pct': round(timings.get('model_fetch_ms', 0) / timings['total_ms'] * 100, 1),
|
||||
'prompt_build_pct': round(timings.get('prompt_build_ms', 0) / timings['total_ms'] * 100, 1),
|
||||
'llm_call_pct': round(timings.get('llm_call_ms', 0) / timings['total_ms'] * 100, 1),
|
||||
}
|
||||
|
||||
if errors:
|
||||
timings['errors'] = errors
|
||||
|
||||
return self.success(data={'test_result': timings})
|
||||
|
||||
|
||||
@group.group_class('executions', '/api/v1/executions')
|
||||
class ExecutionsRouterGroup(group.RouterGroup):
|
||||
"""Workflow execution API router group"""
|
||||
|
||||
async def initialize(self) -> None:
|
||||
# Get all executions (across all workflows)
|
||||
@self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _() -> str:
|
||||
limit = int(quart.request.args.get('limit', 50))
|
||||
offset = int(quart.request.args.get('offset', 0))
|
||||
status = quart.request.args.get('status')
|
||||
executions = await self.ap.workflow_service.get_executions(limit=limit, offset=offset, status=status)
|
||||
return self.success(data=executions)
|
||||
|
||||
# Get single execution
|
||||
@self.route('/<execution_uuid>', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(execution_uuid: str) -> str:
|
||||
execution = await self.ap.workflow_service.get_execution(execution_uuid)
|
||||
if execution is None:
|
||||
return self.http_status(404, -1, 'execution not found')
|
||||
return self.success(data={'execution': execution})
|
||||
|
||||
# Cancel execution
|
||||
@self.route('/<execution_uuid>/cancel', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
|
||||
async def _(execution_uuid: str) -> str:
|
||||
try:
|
||||
await self.ap.workflow_service.cancel_execution(execution_uuid)
|
||||
return self.success()
|
||||
except ValueError as e:
|
||||
return self.http_status(404, -1, str(e))
|
||||
except RuntimeError as e:
|
||||
return self.http_status(400, -1, str(e))
|
||||
@@ -17,6 +17,7 @@ from .groups import platform as groups_platform
|
||||
from .groups import pipelines as groups_pipelines
|
||||
from .groups import knowledge as groups_knowledge
|
||||
from .groups import resources as groups_resources
|
||||
from .groups import workflows as groups_workflows
|
||||
from ...mcp.mount import MCPMount
|
||||
|
||||
importutil.import_modules_in_pkg(groups)
|
||||
@@ -25,6 +26,7 @@ importutil.import_modules_in_pkg(groups_platform)
|
||||
importutil.import_modules_in_pkg(groups_pipelines)
|
||||
importutil.import_modules_in_pkg(groups_knowledge)
|
||||
importutil.import_modules_in_pkg(groups_resources)
|
||||
importutil.import_modules_in_pkg(groups_workflows)
|
||||
|
||||
|
||||
class HTTPController:
|
||||
|
||||
@@ -99,16 +99,23 @@ class BotService:
|
||||
# TODO: 检查配置信息格式
|
||||
bot_data['uuid'] = str(uuid.uuid4())
|
||||
|
||||
# bind the most recently updated pipeline if any exist
|
||||
# Set default binding_type if not provided
|
||||
if 'binding_type' not in bot_data:
|
||||
bot_data['binding_type'] = 'pipeline'
|
||||
|
||||
# checkout the default pipeline (for backward compatibility)
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.select(persistence_pipeline.LegacyPipeline)
|
||||
.order_by(persistence_pipeline.LegacyPipeline.updated_at.desc())
|
||||
.limit(1)
|
||||
sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
|
||||
persistence_pipeline.LegacyPipeline.is_default == True
|
||||
)
|
||||
)
|
||||
pipeline = result.first()
|
||||
if pipeline is not None:
|
||||
bot_data['use_pipeline_uuid'] = pipeline.uuid
|
||||
bot_data['use_pipeline_name'] = pipeline.name
|
||||
# Also set binding_uuid for new unified binding model
|
||||
if 'binding_uuid' not in bot_data:
|
||||
bot_data['binding_uuid'] = pipeline.uuid
|
||||
|
||||
await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(persistence_bot.Bot).values(bot_data))
|
||||
|
||||
@@ -120,26 +127,45 @@ class BotService:
|
||||
|
||||
async def update_bot(self, bot_uuid: str, bot_data: dict) -> None:
|
||||
"""Update bot"""
|
||||
update_data = bot_data.copy()
|
||||
if 'uuid' in bot_data:
|
||||
del bot_data['uuid']
|
||||
|
||||
if 'uuid' in update_data:
|
||||
del update_data['uuid']
|
||||
# Handle binding_type and binding_uuid for the new unified binding model
|
||||
# If binding_type is explicitly set to 'workflow', skip pipeline validation
|
||||
binding_type = bot_data.get('binding_type')
|
||||
|
||||
# set use_pipeline_name
|
||||
if 'use_pipeline_uuid' in update_data:
|
||||
# set use_pipeline_name (for backward compatibility with 'pipeline' binding_type)
|
||||
# Only validate pipeline when binding_type is 'pipeline' or not set (default to pipeline)
|
||||
if 'use_pipeline_uuid' in bot_data and binding_type != 'workflow':
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
|
||||
persistence_pipeline.LegacyPipeline.uuid == update_data['use_pipeline_uuid']
|
||||
persistence_pipeline.LegacyPipeline.uuid == bot_data['use_pipeline_uuid']
|
||||
)
|
||||
)
|
||||
pipeline = result.first()
|
||||
if pipeline is not None:
|
||||
update_data['use_pipeline_name'] = pipeline.name
|
||||
bot_data['use_pipeline_name'] = pipeline.name
|
||||
# Also sync to binding_uuid if binding_type is 'pipeline' or not set
|
||||
if binding_type is None or binding_type == 'pipeline':
|
||||
bot_data['binding_uuid'] = bot_data['use_pipeline_uuid']
|
||||
bot_data['binding_type'] = 'pipeline'
|
||||
else:
|
||||
raise Exception('Pipeline not found')
|
||||
# Only raise error if binding_type is explicitly 'pipeline' or not set
|
||||
if binding_type is None or binding_type == 'pipeline':
|
||||
raise Exception('Pipeline not found')
|
||||
# If binding_type is 'workflow', just clear the use_pipeline_uuid
|
||||
bot_data['use_pipeline_uuid'] = None
|
||||
bot_data['use_pipeline_name'] = None
|
||||
|
||||
# If binding_uuid is set directly (for workflow), clear pipeline fields
|
||||
if 'binding_uuid' in bot_data and binding_type == 'workflow':
|
||||
# For workflow binding, clear pipeline-related fields to avoid confusion
|
||||
bot_data['binding_type'] = 'workflow'
|
||||
bot_data['use_pipeline_uuid'] = None
|
||||
bot_data['use_pipeline_name'] = None
|
||||
|
||||
await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.update(persistence_bot.Bot).values(update_data).where(persistence_bot.Bot.uuid == bot_uuid)
|
||||
sqlalchemy.update(persistence_bot.Bot).values(bot_data).where(persistence_bot.Bot.uuid == bot_uuid)
|
||||
)
|
||||
await self.ap.platform_mgr.remove_bot(bot_uuid)
|
||||
|
||||
|
||||
@@ -243,7 +243,6 @@ class MaintenanceService:
|
||||
tables = {
|
||||
'messages': persistence_monitoring.MonitoringMessage.id,
|
||||
'llm_calls': persistence_monitoring.MonitoringLLMCall.id,
|
||||
'tool_calls': persistence_monitoring.MonitoringToolCall.id,
|
||||
'embedding_calls': persistence_monitoring.MonitoringEmbeddingCall.id,
|
||||
'errors': persistence_monitoring.MonitoringError.id,
|
||||
'sessions': persistence_monitoring.MonitoringSession.session_id,
|
||||
|
||||
@@ -48,17 +48,6 @@ class MCPService:
|
||||
if total_extensions >= max_extensions:
|
||||
raise ValueError(f'Maximum number of extensions ({max_extensions}) reached')
|
||||
|
||||
server_name = str(server_data.get('name') or '').strip()
|
||||
if not server_name:
|
||||
raise ValueError('MCP server name is required')
|
||||
server_data['name'] = server_name
|
||||
|
||||
existing_result = await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.name == server_name)
|
||||
)
|
||||
if existing_result.first() is not None:
|
||||
raise ValueError(f'MCP server already exists: {server_name}')
|
||||
|
||||
server_data['uuid'] = str(uuid.uuid4())
|
||||
await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(persistence_mcp.MCPServer).values(server_data))
|
||||
|
||||
@@ -188,22 +177,10 @@ class MCPService:
|
||||
persisted_session = runtime_mcp_session
|
||||
|
||||
async def _refresh_and_report() -> None:
|
||||
# Testing a persisted server should REUSE its live shared-session
|
||||
# process, not rebuild it. Try a lightweight refresh (a real
|
||||
# list_tools probe over the existing connection) first; only fall
|
||||
# back to a full start() when the session has no live connection
|
||||
# to probe (never connected, or the process is actually gone).
|
||||
needs_start = persisted_session.status == MCPSessionStatus.ERROR or persisted_session.session is None
|
||||
if needs_start:
|
||||
if persisted_session.status == MCPSessionStatus.ERROR:
|
||||
await persisted_session.start()
|
||||
else:
|
||||
try:
|
||||
await persisted_session.refresh()
|
||||
except Exception:
|
||||
# The live connection was stale/dropped: reconnect once
|
||||
# (reusing the live managed process where possible) and
|
||||
# re-probe, instead of reporting a false failure.
|
||||
await persisted_session.start()
|
||||
await persisted_session.refresh()
|
||||
# Surface the discovered tools so the config page can render them
|
||||
# even for an already-hosted server.
|
||||
ctx.metadata['runtime_info'] = persisted_session.get_runtime_info_dict()
|
||||
@@ -244,19 +221,3 @@ class MCPService:
|
||||
context=ctx,
|
||||
)
|
||||
return wrapper.id
|
||||
|
||||
async def get_mcp_server_logs(self, server_name: str, limit: int = 200, level: str | None = None) -> list[dict]:
|
||||
"""Get recent log lines captured from the MCP server's stderr."""
|
||||
session = self.ap.tool_mgr.mcp_tool_loader.get_session(server_name)
|
||||
if not session:
|
||||
return []
|
||||
|
||||
# Get logs from the session's buffer
|
||||
logs = list(session._log_buffer)
|
||||
|
||||
# Filter by level if specified
|
||||
if level:
|
||||
logs = [log for log in logs if log.get('level') == level]
|
||||
|
||||
# Return the most recent 'limit' logs
|
||||
return logs[-limit:]
|
||||
|
||||
@@ -2,7 +2,6 @@ from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
import datetime
|
||||
import json
|
||||
import sqlalchemy
|
||||
|
||||
from ....core import app
|
||||
@@ -51,12 +50,6 @@ class MonitoringService:
|
||||
persistence_monitoring.MonitoringLLMCall.timestamp,
|
||||
persistence_monitoring.MonitoringLLMCall.id,
|
||||
),
|
||||
(
|
||||
'monitoring_tool_calls',
|
||||
persistence_monitoring.MonitoringToolCall,
|
||||
persistence_monitoring.MonitoringToolCall.timestamp,
|
||||
persistence_monitoring.MonitoringToolCall.id,
|
||||
),
|
||||
(
|
||||
'monitoring_embedding_calls',
|
||||
persistence_monitoring.MonitoringEmbeddingCall,
|
||||
@@ -138,68 +131,6 @@ class MonitoringService:
|
||||
await autocommit_conn.execute(sqlalchemy.text('PRAGMA wal_checkpoint(TRUNCATE)'))
|
||||
await autocommit_conn.execute(sqlalchemy.text('VACUUM'))
|
||||
|
||||
def _serialize_tool_payload(self, payload: object, max_length: int = 20000) -> str | None:
|
||||
"""Serialize tool arguments/results for monitoring storage."""
|
||||
if payload is None:
|
||||
return None
|
||||
|
||||
if isinstance(payload, str):
|
||||
text = payload
|
||||
else:
|
||||
try:
|
||||
text = json.dumps(payload, ensure_ascii=False, default=str)
|
||||
except Exception:
|
||||
text = str(payload)
|
||||
|
||||
if len(text) <= max_length:
|
||||
return text
|
||||
|
||||
return f'{text[:max_length]}... [truncated {len(text) - max_length} chars]'
|
||||
|
||||
async def _get_message_for_tool_context(
|
||||
self,
|
||||
message_id: str | None = None,
|
||||
session_id: str | None = None,
|
||||
):
|
||||
if message_id:
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.select(persistence_monitoring.MonitoringMessage).where(
|
||||
persistence_monitoring.MonitoringMessage.id == message_id
|
||||
)
|
||||
)
|
||||
row = result.first()
|
||||
if row:
|
||||
return row[0]
|
||||
|
||||
if not session_id:
|
||||
return None
|
||||
|
||||
user_query = (
|
||||
sqlalchemy.select(persistence_monitoring.MonitoringMessage)
|
||||
.where(
|
||||
sqlalchemy.and_(
|
||||
persistence_monitoring.MonitoringMessage.session_id == session_id,
|
||||
persistence_monitoring.MonitoringMessage.role == 'user',
|
||||
)
|
||||
)
|
||||
.order_by(persistence_monitoring.MonitoringMessage.timestamp.desc())
|
||||
.limit(1)
|
||||
)
|
||||
result = await self.ap.persistence_mgr.execute_async(user_query)
|
||||
row = result.first()
|
||||
if row:
|
||||
return row[0]
|
||||
|
||||
any_query = (
|
||||
sqlalchemy.select(persistence_monitoring.MonitoringMessage)
|
||||
.where(persistence_monitoring.MonitoringMessage.session_id == session_id)
|
||||
.order_by(persistence_monitoring.MonitoringMessage.timestamp.desc())
|
||||
.limit(1)
|
||||
)
|
||||
result = await self.ap.persistence_mgr.execute_async(any_query)
|
||||
row = result.first()
|
||||
return row[0] if row else None
|
||||
|
||||
# ========== Recording Methods ==========
|
||||
|
||||
async def record_message(
|
||||
@@ -289,57 +220,6 @@ class MonitoringService:
|
||||
|
||||
return call_id
|
||||
|
||||
async def record_tool_call(
|
||||
self,
|
||||
tool_name: str,
|
||||
tool_source: str,
|
||||
duration: int,
|
||||
status: str = 'success',
|
||||
bot_id: str | None = None,
|
||||
bot_name: str | None = None,
|
||||
pipeline_id: str | None = None,
|
||||
pipeline_name: str | None = None,
|
||||
session_id: str | None = None,
|
||||
message_id: str | None = None,
|
||||
arguments: object | None = None,
|
||||
result: object | None = None,
|
||||
error_message: str | None = None,
|
||||
) -> str:
|
||||
"""Record a tool call."""
|
||||
context_message = await self._get_message_for_tool_context(message_id=message_id, session_id=session_id)
|
||||
if context_message:
|
||||
bot_id = bot_id or context_message.bot_id
|
||||
bot_name = bot_name or context_message.bot_name
|
||||
pipeline_id = pipeline_id or context_message.pipeline_id
|
||||
pipeline_name = pipeline_name or context_message.pipeline_name
|
||||
session_id = session_id or context_message.session_id
|
||||
message_id = message_id or context_message.id
|
||||
|
||||
call_id = str(uuid.uuid4())
|
||||
call_data = {
|
||||
'id': call_id,
|
||||
'timestamp': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
|
||||
'tool_name': tool_name,
|
||||
'tool_source': tool_source,
|
||||
'duration': max(0, duration),
|
||||
'status': status,
|
||||
'bot_id': bot_id or 'unknown',
|
||||
'bot_name': bot_name or 'Unknown',
|
||||
'pipeline_id': pipeline_id or 'unknown',
|
||||
'pipeline_name': pipeline_name or 'Unknown',
|
||||
'session_id': session_id,
|
||||
'message_id': message_id,
|
||||
'arguments': self._serialize_tool_payload(arguments),
|
||||
'result': self._serialize_tool_payload(result),
|
||||
'error_message': self._serialize_tool_payload(error_message),
|
||||
}
|
||||
|
||||
await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.insert(persistence_monitoring.MonitoringToolCall).values(call_data)
|
||||
)
|
||||
|
||||
return call_id
|
||||
|
||||
async def record_embedding_call(
|
||||
self,
|
||||
model_name: str,
|
||||
@@ -869,58 +749,6 @@ class MonitoringService:
|
||||
total,
|
||||
)
|
||||
|
||||
async def get_tool_calls(
|
||||
self,
|
||||
bot_ids: list[str] | None = None,
|
||||
pipeline_ids: list[str] | None = None,
|
||||
session_ids: list[str] | None = None,
|
||||
start_time: datetime.datetime | None = None,
|
||||
end_time: datetime.datetime | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
) -> tuple[list[dict], int]:
|
||||
"""Get tool calls with filters"""
|
||||
conditions = []
|
||||
|
||||
if bot_ids:
|
||||
conditions.append(persistence_monitoring.MonitoringToolCall.bot_id.in_(bot_ids))
|
||||
if pipeline_ids:
|
||||
conditions.append(persistence_monitoring.MonitoringToolCall.pipeline_id.in_(pipeline_ids))
|
||||
if session_ids:
|
||||
conditions.append(persistence_monitoring.MonitoringToolCall.session_id.in_(session_ids))
|
||||
if start_time:
|
||||
conditions.append(persistence_monitoring.MonitoringToolCall.timestamp >= start_time)
|
||||
if end_time:
|
||||
conditions.append(persistence_monitoring.MonitoringToolCall.timestamp <= end_time)
|
||||
|
||||
count_query = sqlalchemy.select(sqlalchemy.func.count(persistence_monitoring.MonitoringToolCall.id))
|
||||
if conditions:
|
||||
count_query = count_query.where(sqlalchemy.and_(*conditions))
|
||||
|
||||
count_result = await self.ap.persistence_mgr.execute_async(count_query)
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
query = sqlalchemy.select(persistence_monitoring.MonitoringToolCall).order_by(
|
||||
persistence_monitoring.MonitoringToolCall.timestamp.desc()
|
||||
)
|
||||
if conditions:
|
||||
query = query.where(sqlalchemy.and_(*conditions))
|
||||
|
||||
query = query.limit(limit).offset(offset)
|
||||
|
||||
result = await self.ap.persistence_mgr.execute_async(query)
|
||||
tool_calls_rows = result.all()
|
||||
|
||||
return (
|
||||
[
|
||||
self.ap.persistence_mgr.serialize_model(
|
||||
persistence_monitoring.MonitoringToolCall, row[0] if isinstance(row, tuple) else row
|
||||
)
|
||||
for row in tool_calls_rows
|
||||
],
|
||||
total,
|
||||
)
|
||||
|
||||
async def get_embedding_calls(
|
||||
self,
|
||||
start_time: datetime.datetime | None = None,
|
||||
@@ -1143,34 +971,6 @@ class MonitoringService:
|
||||
else:
|
||||
error_llm_calls += 1
|
||||
|
||||
# Get tool calls for this session
|
||||
tool_query = (
|
||||
sqlalchemy.select(persistence_monitoring.MonitoringToolCall)
|
||||
.where(persistence_monitoring.MonitoringToolCall.session_id == session_id)
|
||||
.order_by(persistence_monitoring.MonitoringToolCall.timestamp.asc())
|
||||
)
|
||||
tool_result = await self.ap.persistence_mgr.execute_async(tool_query)
|
||||
tool_rows = tool_result.all()
|
||||
|
||||
tool_calls = [
|
||||
self.ap.persistence_mgr.serialize_model(
|
||||
persistence_monitoring.MonitoringToolCall, row[0] if isinstance(row, tuple) else row
|
||||
)
|
||||
for row in tool_rows
|
||||
]
|
||||
|
||||
total_tool_calls = len(tool_rows)
|
||||
success_tool_calls = 0
|
||||
error_tool_calls = 0
|
||||
total_tool_duration = 0
|
||||
for row in tool_rows:
|
||||
tool_call = row[0] if isinstance(row, tuple) else row
|
||||
total_tool_duration += tool_call.duration
|
||||
if tool_call.status == 'success':
|
||||
success_tool_calls += 1
|
||||
else:
|
||||
error_tool_calls += 1
|
||||
|
||||
# Get errors for this session
|
||||
error_query = (
|
||||
sqlalchemy.select(persistence_monitoring.MonitoringError)
|
||||
@@ -1214,14 +1014,6 @@ class MonitoringService:
|
||||
'total_tokens': total_tokens,
|
||||
'average_duration_ms': int(total_duration / total_llm_calls) if total_llm_calls > 0 else 0,
|
||||
},
|
||||
'tool_calls': tool_calls,
|
||||
'tool_stats': {
|
||||
'total_calls': total_tool_calls,
|
||||
'success_calls': success_tool_calls,
|
||||
'error_calls': error_tool_calls,
|
||||
'total_duration_ms': total_tool_duration,
|
||||
'average_duration_ms': int(total_tool_duration / total_tool_calls) if total_tool_calls > 0 else 0,
|
||||
},
|
||||
'errors': errors,
|
||||
'session_duration_seconds': session_duration_seconds,
|
||||
}
|
||||
|
||||
@@ -73,6 +73,20 @@ class PipelineService:
|
||||
|
||||
return self.ap.persistence_mgr.serialize_model(persistence_pipeline.LegacyPipeline, pipeline)
|
||||
|
||||
async def get_pipeline_by_name(self, pipeline_name: str) -> dict | None:
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
|
||||
persistence_pipeline.LegacyPipeline.name == pipeline_name
|
||||
)
|
||||
)
|
||||
|
||||
pipeline = result.first()
|
||||
|
||||
if pipeline is None:
|
||||
return None
|
||||
|
||||
return self.ap.persistence_mgr.serialize_model(persistence_pipeline.LegacyPipeline, pipeline)
|
||||
|
||||
async def create_pipeline(self, pipeline_data: dict, default: bool = False) -> str:
|
||||
from ....utils import paths as path_utils
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -32,6 +32,7 @@ from ..api.http.service import mcp as mcp_service
|
||||
from ..api.http.service import apikey as apikey_service
|
||||
from ..api.http.service import webhook as webhook_service
|
||||
from ..api.http.service import monitoring as monitoring_service
|
||||
from ..api.http.service import workflow as workflow_service
|
||||
from ..api.http.service import skill as skill_service
|
||||
from ..api.http.service import maintenance as maintenance_service
|
||||
from ..discover import engine as discover_engine
|
||||
@@ -153,6 +154,8 @@ class Application:
|
||||
|
||||
webhook_service: webhook_service.WebhookService = None
|
||||
|
||||
workflow_service: workflow_service.WorkflowService = None
|
||||
|
||||
telemetry: telemetry_module.TelemetryManager = None
|
||||
|
||||
survey: survey_module.SurveyManager = None
|
||||
@@ -255,6 +258,22 @@ class Application:
|
||||
scopes=[core_entities.LifecycleControlScope.APPLICATION],
|
||||
)
|
||||
|
||||
async def workflow_execution_cleanup_loop():
|
||||
check_interval_seconds = 60
|
||||
while True:
|
||||
try:
|
||||
cancelled = await self.workflow_service.cleanup_stale_executions()
|
||||
if cancelled > 0:
|
||||
self.logger.info(f'Workflow execution auto-cleanup: cancelled {cancelled} stale executions')
|
||||
except Exception as e:
|
||||
self.logger.warning(f'Workflow execution auto-cleanup error: {e}')
|
||||
await asyncio.sleep(check_interval_seconds)
|
||||
|
||||
self.task_mgr.create_task(
|
||||
workflow_execution_cleanup_loop(),
|
||||
name='workflow-execution-cleanup',
|
||||
scopes=[core_entities.LifecycleControlScope.APPLICATION],
|
||||
)
|
||||
# Start storage/log maintenance task if enabled
|
||||
storage_cleanup_cfg = self.instance_config.data.get('storage', {}).get('cleanup', {})
|
||||
if storage_cleanup_cfg.get('enabled', True) and self.maintenance_service is not None:
|
||||
|
||||
@@ -29,6 +29,7 @@ from ...api.http.service import mcp as mcp_service
|
||||
from ...api.http.service import apikey as apikey_service
|
||||
from ...api.http.service import webhook as webhook_service
|
||||
from ...api.http.service import monitoring as monitoring_service
|
||||
from ...api.http.service import workflow as workflow_service
|
||||
from ...api.http.service import skill as skill_service
|
||||
from ...skill import manager as skill_mgr
|
||||
from ...api.http.service import maintenance as maintenance_service
|
||||
@@ -89,6 +90,9 @@ class BuildAppStage(stage.BootingStage):
|
||||
webhook_service_inst = webhook_service.WebhookService(ap)
|
||||
ap.webhook_service = webhook_service_inst
|
||||
|
||||
workflow_service_inst = workflow_service.WorkflowService(ap)
|
||||
ap.workflow_service = workflow_service_inst
|
||||
|
||||
skill_service_inst = skill_service.SkillService(ap)
|
||||
ap.skill_service = skill_service_inst
|
||||
|
||||
|
||||
@@ -231,3 +231,34 @@ class LoadConfigStage(stage.BootingStage):
|
||||
ap.pipeline_config_meta_safety = await load_resource_yaml_template_data('metadata/pipeline/safety.yaml')
|
||||
ap.pipeline_config_meta_ai = await load_resource_yaml_template_data('metadata/pipeline/ai.yaml')
|
||||
ap.pipeline_config_meta_output = await load_resource_yaml_template_data('metadata/pipeline/output.yaml')
|
||||
|
||||
# Load workflow node metadata from YAML files. YAML is the source of
|
||||
# truth for workflow editor metadata; Python classes provide execution
|
||||
# logic and are bound through the registry.
|
||||
from langbot.pkg.workflow.metadata import NodeMetadataLoader
|
||||
from langbot.pkg.workflow.registry import NodeTypeRegistry
|
||||
|
||||
workflow_metadata_loader = NodeMetadataLoader()
|
||||
workflow_node_count = await workflow_metadata_loader.load_core_metadata()
|
||||
ap.workflow_node_configs = workflow_metadata_loader.get_all_metadata()
|
||||
ap.workflow_node_metadata_loader = workflow_metadata_loader
|
||||
|
||||
workflow_registry = NodeTypeRegistry.instance()
|
||||
for node_config in ap.workflow_node_configs.values():
|
||||
workflow_registry.register_metadata(node_config, source=node_config.get('_source', 'core'))
|
||||
|
||||
# Auto-discover and register workflow nodes using discovery engine
|
||||
if hasattr(ap, 'discover') and ap.discover is not None:
|
||||
workflow_registry.discover_nodes(ap.discover)
|
||||
|
||||
workflow_load_errors = workflow_metadata_loader.get_load_errors()
|
||||
if workflow_load_errors:
|
||||
print(f'Workflow node metadata load errors: {len(workflow_load_errors)}')
|
||||
for error in workflow_load_errors:
|
||||
print(f" - {error.get('file')}: {error.get('error')}")
|
||||
|
||||
print(
|
||||
f'Loaded {workflow_node_count} workflow node metadata files; '
|
||||
f'registered {workflow_registry.metadata_count()} metadata definitions, '
|
||||
f'{workflow_registry.count()} node types'
|
||||
)
|
||||
|
||||
@@ -304,3 +304,65 @@ class ComponentDiscoveryEngine:
|
||||
if component.kind == kind:
|
||||
result.append(component)
|
||||
return result
|
||||
|
||||
def discover_workflow_nodes(self, nodes_dir: str) -> typing.List[typing.Type]:
|
||||
"""Discover workflow node classes from a directory of Python modules.
|
||||
|
||||
Scans all .py files in the given directory, imports them, and collects
|
||||
classes that are subclasses of WorkflowNode.
|
||||
|
||||
Args:
|
||||
nodes_dir: Directory path like 'pkg/workflow/nodes/'
|
||||
|
||||
Returns:
|
||||
List of WorkflowNode subclasses found
|
||||
"""
|
||||
from langbot.pkg.workflow.node import WorkflowNode
|
||||
|
||||
node_classes: typing.List[typing.Type[WorkflowNode]] = []
|
||||
|
||||
# Normalize path
|
||||
if nodes_dir.endswith('/'):
|
||||
nodes_dir = nodes_dir[:-1]
|
||||
|
||||
# Import the nodes package to trigger all module imports
|
||||
module_path = nodes_dir.replace('/', '.').replace('\\', '.')
|
||||
package_path = module_path
|
||||
|
||||
try:
|
||||
# Import the package __init__ to trigger submodule imports
|
||||
importlib.import_module(f'langbot.{package_path}')
|
||||
except ImportError:
|
||||
self.ap.logger.warning(f'Failed to import workflow nodes package: langbot.{package_path}')
|
||||
|
||||
# Since workflow/__init__.py is empty, explicitly import all .py files in the nodes directory
|
||||
import os
|
||||
# engine.py is in langbot/pkg/discover/, nodes are in langbot/pkg/workflow/nodes/
|
||||
nodes_abs_path = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'workflow', 'nodes'))
|
||||
if os.path.isdir(nodes_abs_path):
|
||||
for filename in os.listdir(nodes_abs_path):
|
||||
if filename.endswith('.py') and not filename.startswith('_'):
|
||||
module_name = filename[:-3]
|
||||
try:
|
||||
importlib.import_module(f'langbot.{package_path}.{module_name}')
|
||||
except ImportError as e:
|
||||
self.ap.logger.warning(f'Failed to import workflow node module: {module_name}: {e}')
|
||||
|
||||
# Now collect all WorkflowNode subclasses from sys.modules
|
||||
import sys
|
||||
prefix = f'langbot.{package_path}.'
|
||||
for mod_name, mod in sys.modules.items():
|
||||
if mod_name.startswith(prefix) and mod is not None:
|
||||
for attr_name in dir(mod):
|
||||
attr = getattr(mod, attr_name)
|
||||
if (
|
||||
isinstance(attr, type)
|
||||
and issubclass(attr, WorkflowNode)
|
||||
and attr is not WorkflowNode
|
||||
and hasattr(attr, 'type_name')
|
||||
and attr.type_name
|
||||
):
|
||||
if attr not in node_classes:
|
||||
node_classes.append(attr)
|
||||
|
||||
return node_classes
|
||||
|
||||
@@ -31,6 +31,13 @@ class Bot(Base):
|
||||
use_pipeline_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=True)
|
||||
use_pipeline_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=True)
|
||||
pipeline_routing_rules = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, server_default='[]')
|
||||
|
||||
# New unified binding fields
|
||||
# binding_type: 'pipeline' or 'workflow'
|
||||
binding_type = sqlalchemy.Column(sqlalchemy.String(32), nullable=False, server_default='pipeline')
|
||||
# binding_uuid: UUID of the bound Pipeline or Workflow
|
||||
binding_uuid = sqlalchemy.Column(sqlalchemy.String(64), nullable=True)
|
||||
|
||||
created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
|
||||
updated_at = sqlalchemy.Column(
|
||||
sqlalchemy.DateTime,
|
||||
|
||||
@@ -49,28 +49,6 @@ class MonitoringLLMCall(Base):
|
||||
message_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True, index=True) # Associated message ID
|
||||
|
||||
|
||||
class MonitoringToolCall(Base):
|
||||
"""Tool call records"""
|
||||
|
||||
__tablename__ = 'monitoring_tool_calls'
|
||||
|
||||
id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
|
||||
timestamp = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, index=True)
|
||||
tool_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
|
||||
tool_source = sqlalchemy.Column(sqlalchemy.String(50), nullable=False) # native, plugin, mcp, skill
|
||||
duration = sqlalchemy.Column(sqlalchemy.Integer, nullable=False) # milliseconds
|
||||
status = sqlalchemy.Column(sqlalchemy.String(50), nullable=False) # success, error
|
||||
bot_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
|
||||
bot_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
|
||||
pipeline_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
|
||||
pipeline_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
|
||||
session_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True, index=True)
|
||||
message_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True, index=True)
|
||||
arguments = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
|
||||
result = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
|
||||
error_message = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
|
||||
|
||||
|
||||
class MonitoringSession(Base):
|
||||
"""Session tracking records"""
|
||||
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
"""Workflow persistence entities"""
|
||||
|
||||
import sqlalchemy
|
||||
|
||||
from .base import Base
|
||||
|
||||
|
||||
class Workflow(Base):
|
||||
"""Workflow definition"""
|
||||
|
||||
__tablename__ = 'workflows'
|
||||
|
||||
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
|
||||
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
|
||||
description = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
|
||||
emoji = sqlalchemy.Column(sqlalchemy.String(10), nullable=True, default='🔄')
|
||||
version = sqlalchemy.Column(sqlalchemy.Integer, nullable=False, default=1)
|
||||
is_enabled = sqlalchemy.Column(sqlalchemy.Boolean, nullable=False, default=True)
|
||||
|
||||
# Workflow definition stored as JSON
|
||||
# Contains: nodes, edges, variables, settings
|
||||
definition = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, default={})
|
||||
|
||||
# Global config (inherited from Pipeline capabilities)
|
||||
# Contains: safety, output configs
|
||||
global_config = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, default={})
|
||||
|
||||
# Extensions preferences (same as Pipeline)
|
||||
extensions_preferences = sqlalchemy.Column(
|
||||
sqlalchemy.JSON,
|
||||
nullable=False,
|
||||
default={'enable_all_plugins': True, 'enable_all_mcp_servers': True, 'plugins': [], 'mcp_servers': []},
|
||||
)
|
||||
|
||||
created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
|
||||
updated_at = sqlalchemy.Column(
|
||||
sqlalchemy.DateTime,
|
||||
nullable=False,
|
||||
server_default=sqlalchemy.func.now(),
|
||||
onupdate=sqlalchemy.func.now(),
|
||||
)
|
||||
|
||||
|
||||
class WorkflowVersion(Base):
|
||||
"""Workflow version history"""
|
||||
|
||||
__tablename__ = 'workflow_versions'
|
||||
|
||||
id = sqlalchemy.Column(sqlalchemy.Integer, primary_key=True, autoincrement=True)
|
||||
workflow_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
|
||||
version = sqlalchemy.Column(sqlalchemy.Integer, nullable=False)
|
||||
definition = sqlalchemy.Column(sqlalchemy.JSON, nullable=False)
|
||||
global_config = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, default={})
|
||||
created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
|
||||
created_by = sqlalchemy.Column(sqlalchemy.String(255), nullable=True)
|
||||
|
||||
__table_args__ = (sqlalchemy.UniqueConstraint('workflow_uuid', 'version', name='uq_workflow_version'),)
|
||||
|
||||
|
||||
class WorkflowTrigger(Base):
|
||||
"""Workflow trigger configuration"""
|
||||
|
||||
__tablename__ = 'workflow_triggers'
|
||||
|
||||
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
|
||||
workflow_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
|
||||
type = sqlalchemy.Column(sqlalchemy.String(50), nullable=False) # message, cron, event, webhook
|
||||
config = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, default={})
|
||||
is_enabled = sqlalchemy.Column(sqlalchemy.Boolean, nullable=False, default=True)
|
||||
priority = sqlalchemy.Column(sqlalchemy.Integer, nullable=False, default=0)
|
||||
created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
|
||||
updated_at = sqlalchemy.Column(
|
||||
sqlalchemy.DateTime,
|
||||
nullable=False,
|
||||
server_default=sqlalchemy.func.now(),
|
||||
onupdate=sqlalchemy.func.now(),
|
||||
)
|
||||
|
||||
|
||||
class WorkflowExecution(Base):
|
||||
"""Workflow execution record"""
|
||||
|
||||
__tablename__ = 'workflow_executions'
|
||||
|
||||
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
|
||||
workflow_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
|
||||
workflow_version = sqlalchemy.Column(sqlalchemy.Integer, nullable=False)
|
||||
status = sqlalchemy.Column(sqlalchemy.String(20), nullable=False) # pending, running, completed, failed, cancelled
|
||||
trigger_type = sqlalchemy.Column(sqlalchemy.String(50), nullable=True)
|
||||
trigger_data = sqlalchemy.Column(sqlalchemy.JSON, nullable=True)
|
||||
variables = sqlalchemy.Column(sqlalchemy.JSON, nullable=True)
|
||||
start_time = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
|
||||
end_time = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
|
||||
error = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
|
||||
created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
|
||||
|
||||
|
||||
class WorkflowNodeExecution(Base):
|
||||
"""Workflow node execution record"""
|
||||
|
||||
__tablename__ = 'workflow_node_executions'
|
||||
|
||||
id = sqlalchemy.Column(sqlalchemy.Integer, primary_key=True, autoincrement=True)
|
||||
execution_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
|
||||
node_id = sqlalchemy.Column(sqlalchemy.String(100), nullable=False)
|
||||
node_type = sqlalchemy.Column(sqlalchemy.String(50), nullable=False)
|
||||
status = sqlalchemy.Column(sqlalchemy.String(20), nullable=False) # pending, running, completed, failed, skipped
|
||||
inputs = sqlalchemy.Column(sqlalchemy.JSON, nullable=True)
|
||||
outputs = sqlalchemy.Column(sqlalchemy.JSON, nullable=True)
|
||||
start_time = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
|
||||
end_time = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
|
||||
error = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
|
||||
retry_count = sqlalchemy.Column(sqlalchemy.Integer, nullable=False, default=0)
|
||||
|
||||
|
||||
class ScheduledJob(Base):
|
||||
"""Scheduled job for cron triggers"""
|
||||
|
||||
__tablename__ = 'workflow_scheduled_jobs'
|
||||
|
||||
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
|
||||
trigger_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
|
||||
cron_expression = sqlalchemy.Column(sqlalchemy.String(100), nullable=True)
|
||||
next_run_time = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
|
||||
last_run_time = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
|
||||
is_enabled = sqlalchemy.Column(sqlalchemy.Boolean, nullable=False, default=True)
|
||||
@@ -0,0 +1,207 @@
|
||||
"""add workflow tables and bot binding fields
|
||||
|
||||
Revision ID: 0009_add_workflow_tables
|
||||
Revises: 0008_mcp_resource_prefs
|
||||
Create Date: 2026-07-01
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = '0009_add_workflow_tables'
|
||||
down_revision = '0008_mcp_resource_prefs'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _table_exists(conn: sa.Connection, table_name: str) -> bool:
|
||||
return table_name in sa.inspect(conn).get_table_names()
|
||||
|
||||
|
||||
def _has_column(conn: sa.Connection, table_name: str, column_name: str) -> bool:
|
||||
if not _table_exists(conn, table_name):
|
||||
return False
|
||||
return column_name in {column['name'] for column in sa.inspect(conn).get_columns(table_name)}
|
||||
|
||||
|
||||
def _has_index_for_columns(conn: sa.Connection, table_name: str, columns: tuple[str, ...]) -> bool:
|
||||
if not _table_exists(conn, table_name):
|
||||
return False
|
||||
for index in sa.inspect(conn).get_indexes(table_name):
|
||||
if tuple(index.get('column_names') or ()) == columns:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _ensure_index(conn: sa.Connection, table_name: str, index_name: str, columns: list[str]) -> None:
|
||||
if _has_index_for_columns(conn, table_name, tuple(columns)):
|
||||
return
|
||||
op.create_index(index_name, table_name, columns)
|
||||
|
||||
|
||||
def _create_workflow_tables(conn: sa.Connection) -> None:
|
||||
if not _table_exists(conn, 'workflows'):
|
||||
op.create_table(
|
||||
'workflows',
|
||||
sa.Column('uuid', sa.String(255), primary_key=True),
|
||||
sa.Column('name', sa.String(255), nullable=False),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('emoji', sa.String(10), nullable=True),
|
||||
sa.Column('version', sa.Integer(), nullable=False, server_default='1'),
|
||||
sa.Column('is_enabled', sa.Boolean(), nullable=False, server_default=sa.true()),
|
||||
sa.Column('definition', sa.JSON(), nullable=False, server_default=sa.text("'{}'")),
|
||||
sa.Column('global_config', sa.JSON(), nullable=False, server_default=sa.text("'{}'")),
|
||||
sa.Column(
|
||||
'extensions_preferences',
|
||||
sa.JSON(),
|
||||
nullable=False,
|
||||
server_default=sa.text(
|
||||
'\'{"enable_all_plugins": true, "enable_all_mcp_servers": true, "plugins": [], "mcp_servers": []}\''
|
||||
),
|
||||
),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column('updated_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
if not _table_exists(conn, 'workflow_versions'):
|
||||
op.create_table(
|
||||
'workflow_versions',
|
||||
sa.Column('id', sa.Integer(), primary_key=True, autoincrement=True),
|
||||
sa.Column('workflow_uuid', sa.String(255), nullable=False),
|
||||
sa.Column('version', sa.Integer(), nullable=False),
|
||||
sa.Column('definition', sa.JSON(), nullable=False),
|
||||
sa.Column('global_config', sa.JSON(), nullable=False, server_default=sa.text("'{}'")),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column('created_by', sa.String(255), nullable=True),
|
||||
sa.UniqueConstraint('workflow_uuid', 'version', name='uq_workflow_version'),
|
||||
)
|
||||
|
||||
if not _table_exists(conn, 'workflow_triggers'):
|
||||
op.create_table(
|
||||
'workflow_triggers',
|
||||
sa.Column('uuid', sa.String(255), primary_key=True),
|
||||
sa.Column('workflow_uuid', sa.String(255), nullable=False),
|
||||
sa.Column('type', sa.String(50), nullable=False),
|
||||
sa.Column('config', sa.JSON(), nullable=False, server_default=sa.text("'{}'")),
|
||||
sa.Column('is_enabled', sa.Boolean(), nullable=False, server_default=sa.true()),
|
||||
sa.Column('priority', sa.Integer(), nullable=False, server_default='0'),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column('updated_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
if not _table_exists(conn, 'workflow_executions'):
|
||||
op.create_table(
|
||||
'workflow_executions',
|
||||
sa.Column('uuid', sa.String(255), primary_key=True),
|
||||
sa.Column('workflow_uuid', sa.String(255), nullable=False),
|
||||
sa.Column('workflow_version', sa.Integer(), nullable=False),
|
||||
sa.Column('status', sa.String(20), nullable=False),
|
||||
sa.Column('trigger_type', sa.String(50), nullable=True),
|
||||
sa.Column('trigger_data', sa.JSON(), nullable=True),
|
||||
sa.Column('variables', sa.JSON(), nullable=True),
|
||||
sa.Column('start_time', sa.DateTime(), nullable=True),
|
||||
sa.Column('end_time', sa.DateTime(), nullable=True),
|
||||
sa.Column('error', sa.Text(), nullable=True),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
if not _table_exists(conn, 'workflow_node_executions'):
|
||||
op.create_table(
|
||||
'workflow_node_executions',
|
||||
sa.Column('id', sa.Integer(), primary_key=True, autoincrement=True),
|
||||
sa.Column('execution_uuid', sa.String(255), nullable=False),
|
||||
sa.Column('node_id', sa.String(100), nullable=False),
|
||||
sa.Column('node_type', sa.String(50), nullable=False),
|
||||
sa.Column('status', sa.String(20), nullable=False),
|
||||
sa.Column('inputs', sa.JSON(), nullable=True),
|
||||
sa.Column('outputs', sa.JSON(), nullable=True),
|
||||
sa.Column('start_time', sa.DateTime(), nullable=True),
|
||||
sa.Column('end_time', sa.DateTime(), nullable=True),
|
||||
sa.Column('error', sa.Text(), nullable=True),
|
||||
sa.Column('retry_count', sa.Integer(), nullable=False, server_default='0'),
|
||||
)
|
||||
|
||||
if not _table_exists(conn, 'workflow_scheduled_jobs'):
|
||||
op.create_table(
|
||||
'workflow_scheduled_jobs',
|
||||
sa.Column('uuid', sa.String(255), primary_key=True),
|
||||
sa.Column('trigger_uuid', sa.String(255), nullable=False),
|
||||
sa.Column('cron_expression', sa.String(100), nullable=True),
|
||||
sa.Column('next_run_time', sa.DateTime(), nullable=True),
|
||||
sa.Column('last_run_time', sa.DateTime(), nullable=True),
|
||||
sa.Column('is_enabled', sa.Boolean(), nullable=False, server_default=sa.true()),
|
||||
)
|
||||
|
||||
_ensure_index(conn, 'workflow_versions', 'ix_workflow_versions_workflow_uuid', ['workflow_uuid'])
|
||||
_ensure_index(conn, 'workflow_triggers', 'ix_workflow_triggers_workflow_uuid', ['workflow_uuid'])
|
||||
_ensure_index(conn, 'workflow_executions', 'ix_workflow_executions_workflow_uuid', ['workflow_uuid'])
|
||||
_ensure_index(
|
||||
conn,
|
||||
'workflow_node_executions',
|
||||
'ix_workflow_node_executions_execution_uuid',
|
||||
['execution_uuid'],
|
||||
)
|
||||
_ensure_index(conn, 'workflow_scheduled_jobs', 'ix_workflow_scheduled_jobs_trigger_uuid', ['trigger_uuid'])
|
||||
|
||||
|
||||
def _add_bot_binding_fields(conn: sa.Connection) -> None:
|
||||
if not _table_exists(conn, 'bots'):
|
||||
return
|
||||
|
||||
if not _has_column(conn, 'bots', 'binding_type'):
|
||||
op.add_column(
|
||||
'bots',
|
||||
sa.Column('binding_type', sa.String(32), nullable=False, server_default='pipeline'),
|
||||
)
|
||||
|
||||
if not _has_column(conn, 'bots', 'binding_uuid'):
|
||||
op.add_column('bots', sa.Column('binding_uuid', sa.String(64), nullable=True))
|
||||
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE bots
|
||||
SET binding_uuid = use_pipeline_uuid
|
||||
WHERE use_pipeline_uuid IS NOT NULL
|
||||
AND use_pipeline_uuid != ''
|
||||
AND (binding_uuid IS NULL OR binding_uuid = '')
|
||||
""")
|
||||
)
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE bots
|
||||
SET binding_type = 'pipeline'
|
||||
WHERE binding_uuid IS NOT NULL
|
||||
AND binding_uuid != ''
|
||||
AND (binding_type IS NULL OR binding_type = '')
|
||||
""")
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
_create_workflow_tables(conn)
|
||||
_add_bot_binding_fields(conn)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
if _has_column(conn, 'bots', 'binding_uuid'):
|
||||
with op.batch_alter_table('bots') as batch_op:
|
||||
batch_op.drop_column('binding_uuid')
|
||||
if _has_column(conn, 'bots', 'binding_type'):
|
||||
with op.batch_alter_table('bots') as batch_op:
|
||||
batch_op.drop_column('binding_type')
|
||||
|
||||
for table_name in (
|
||||
'workflow_scheduled_jobs',
|
||||
'workflow_node_executions',
|
||||
'workflow_executions',
|
||||
'workflow_triggers',
|
||||
'workflow_versions',
|
||||
'workflows',
|
||||
):
|
||||
if _table_exists(conn, table_name):
|
||||
op.drop_table(table_name)
|
||||
@@ -1,17 +0,0 @@
|
||||
from langbot.pkg.entity.persistence import monitoring as persistence_monitoring
|
||||
from .. import migration
|
||||
|
||||
|
||||
@migration.migration_class(26)
|
||||
class DBMigrateMonitoringToolCalls(migration.DBMigration):
|
||||
"""Add monitoring_tool_calls table"""
|
||||
|
||||
async def upgrade(self):
|
||||
"""Upgrade"""
|
||||
async with self.ap.persistence_mgr.get_db_engine().begin() as conn:
|
||||
await conn.run_sync(persistence_monitoring.MonitoringToolCall.__table__.create, checkfirst=True)
|
||||
|
||||
async def downgrade(self):
|
||||
"""Downgrade"""
|
||||
async with self.ap.persistence_mgr.get_db_engine().begin() as conn:
|
||||
await conn.run_sync(persistence_monitoring.MonitoringToolCall.__table__.drop, checkfirst=True)
|
||||
@@ -13,7 +13,7 @@ import langbot_plugin.api.entities.builtin.platform.message as platform_message
|
||||
import langbot_plugin.api.entities.builtin.platform.events as platform_events
|
||||
import langbot_plugin.api.entities.events as events
|
||||
from ..utils import importutil
|
||||
from .config_coercion import coerce_pipeline_config
|
||||
from .config import coerce_pipeline_config
|
||||
|
||||
import langbot_plugin.api.entities.builtin.provider.session as provider_session
|
||||
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
|
||||
@@ -168,7 +168,7 @@ class RuntimePipeline:
|
||||
bot_message=query.resp_messages[-1],
|
||||
message=result.user_notice,
|
||||
quote_origin=query.pipeline_config['output']['misc']['quote-origin'],
|
||||
is_final=[msg.is_final for msg in query.resp_messages][-1],
|
||||
is_final=[msg.is_final for msg in query.resp_messages][0],
|
||||
)
|
||||
else:
|
||||
await query.adapter.reply_message(
|
||||
@@ -295,9 +295,9 @@ class RuntimePipeline:
|
||||
# Record query start and store message_id
|
||||
message_id = ''
|
||||
try:
|
||||
from . import monitoring_helper
|
||||
from . import monitor
|
||||
|
||||
message_id = await monitoring_helper.MonitoringHelper.record_query_start(
|
||||
message_id = await monitor.MonitoringHelper.record_query_start(
|
||||
ap=self.ap,
|
||||
query=query,
|
||||
bot_id=query.bot_uuid or 'unknown',
|
||||
@@ -349,7 +349,7 @@ class RuntimePipeline:
|
||||
# Record query success only if no error occurred during processing
|
||||
if not query.variables.get('_monitoring_has_error', False):
|
||||
try:
|
||||
await monitoring_helper.MonitoringHelper.record_query_success(
|
||||
await monitor.MonitoringHelper.record_query_success(
|
||||
ap=self.ap,
|
||||
message_id=message_id,
|
||||
query=query,
|
||||
@@ -359,7 +359,7 @@ class RuntimePipeline:
|
||||
|
||||
# Record bot response message
|
||||
try:
|
||||
await monitoring_helper.MonitoringHelper.record_query_response(
|
||||
await monitor.MonitoringHelper.record_query_response(
|
||||
ap=self.ap,
|
||||
query=query,
|
||||
bot_id=query.bot_uuid or 'unknown',
|
||||
@@ -378,9 +378,9 @@ class RuntimePipeline:
|
||||
|
||||
# Record query error
|
||||
try:
|
||||
from . import monitoring_helper
|
||||
from . import monitor
|
||||
|
||||
await monitoring_helper.MonitoringHelper.record_query_error(
|
||||
await monitor.MonitoringHelper.record_query_error(
|
||||
ap=self.ap,
|
||||
query=query,
|
||||
bot_id=query.bot_uuid or 'unknown',
|
||||
@@ -395,7 +395,8 @@ class RuntimePipeline:
|
||||
|
||||
finally:
|
||||
self.ap.logger.debug(f'Query {query.query_id} processed')
|
||||
del self.ap.query_pool.cached_queries[query.query_id]
|
||||
# Use pop with default to avoid KeyError if query was never cached
|
||||
self.ap.query_pool.cached_queries.pop(query.query_id, None)
|
||||
|
||||
|
||||
class PipelineManager:
|
||||
|
||||
@@ -42,13 +42,9 @@ class QueryPool:
|
||||
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
|
||||
pipeline_uuid: typing.Optional[str] = None,
|
||||
routed_by_rule: bool = False,
|
||||
variables: typing.Optional[dict[str, typing.Any]] = None,
|
||||
) -> pipeline_query.Query:
|
||||
async with self.condition:
|
||||
query_id = self.query_id_counter
|
||||
initial_variables: dict[str, typing.Any] = {'_routed_by_rule': routed_by_rule}
|
||||
if variables:
|
||||
initial_variables.update(variables)
|
||||
query = pipeline_query.Query(
|
||||
bot_uuid=bot_uuid,
|
||||
query_id=query_id,
|
||||
@@ -57,7 +53,7 @@ class QueryPool:
|
||||
sender_id=sender_id,
|
||||
message_event=message_event,
|
||||
message_chain=message_chain,
|
||||
variables=initial_variables,
|
||||
variables={'_routed_by_rule': routed_by_rule},
|
||||
resp_messages=[],
|
||||
resp_message_chain=[],
|
||||
adapter=adapter,
|
||||
|
||||
@@ -45,7 +45,7 @@ class SendResponseBackStage(stage.PipelineStage):
|
||||
|
||||
try:
|
||||
if await query.adapter.is_stream_output_supported() and has_chunks:
|
||||
is_final = [msg.is_final for msg in query.resp_messages][-1]
|
||||
is_final = [msg.is_final for msg in query.resp_messages][0]
|
||||
await query.adapter.reply_message_chunk(
|
||||
message_source=query.message_event,
|
||||
bot_message=query.resp_messages[-1],
|
||||
|
||||
@@ -2,12 +2,14 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import logging
|
||||
import traceback
|
||||
import sqlalchemy
|
||||
|
||||
from ..core import app, entities as core_entities, taskmgr
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from ..discover import engine
|
||||
|
||||
from ..entity.persistence import bot as persistence_bot
|
||||
@@ -54,29 +56,24 @@ class RuntimeBot:
|
||||
self.task_context = taskmgr.TaskContext()
|
||||
self.logger = logger
|
||||
|
||||
@staticmethod
|
||||
def _match_operator(actual: str, operator: str, expected: str) -> bool:
|
||||
"""Evaluate a single operator condition."""
|
||||
if operator == 'eq':
|
||||
return actual == expected
|
||||
elif operator == 'neq':
|
||||
return actual != expected
|
||||
elif operator == 'contains':
|
||||
return expected in actual
|
||||
elif operator == 'not_contains':
|
||||
return expected not in actual
|
||||
elif operator == 'starts_with':
|
||||
return actual.startswith(expected)
|
||||
elif operator == 'regex':
|
||||
try:
|
||||
return bool(re.search(expected, actual))
|
||||
except re.error:
|
||||
return False
|
||||
return False
|
||||
|
||||
PIPELINE_DISCARD = '__discard__'
|
||||
PIPELINE_DISCARD_DISPLAY_NAME = 'Discarded'
|
||||
|
||||
def get_binding_info(self) -> tuple[str, str | None]:
|
||||
"""Get the binding type and UUID for this bot.
|
||||
|
||||
Returns:
|
||||
tuple: (binding_type, binding_uuid) where binding_type is 'pipeline' or 'workflow'
|
||||
"""
|
||||
binding_type = getattr(self.bot_entity, 'binding_type', 'pipeline') or 'pipeline'
|
||||
binding_uuid = getattr(self.bot_entity, 'binding_uuid', None)
|
||||
|
||||
# Fallback to use_pipeline_uuid for backward compatibility
|
||||
if not binding_uuid and binding_type == 'pipeline':
|
||||
binding_uuid = self.bot_entity.use_pipeline_uuid
|
||||
|
||||
return binding_type, binding_uuid
|
||||
|
||||
def resolve_pipeline_uuid(
|
||||
self,
|
||||
launcher_type: str,
|
||||
@@ -84,56 +81,94 @@ class RuntimeBot:
|
||||
message_text: str,
|
||||
message_element_types: list[str] | None = None,
|
||||
) -> tuple[str | None, bool]:
|
||||
"""Resolve pipeline UUID based on routing rules.
|
||||
"""Resolve pipeline UUID for message processing.
|
||||
|
||||
Rules are evaluated in order; first match wins.
|
||||
Falls back to use_pipeline_uuid if no rule matches.
|
||||
|
||||
Rule types:
|
||||
- launcher_type: session type ("person" / "group")
|
||||
- launcher_id: session / group id
|
||||
- message_content: message text content
|
||||
- message_has_element: message contains element of given type
|
||||
(Image, Voice, File, Forward, Face, At, AtAll, Quote)
|
||||
Operators: eq (has), neq (doesn't have)
|
||||
|
||||
Operators: eq, neq, contains, not_contains, starts_with, regex
|
||||
|
||||
When pipeline_uuid is ``__discard__``, the message should be
|
||||
silently dropped by the caller.
|
||||
NOTE: Routing rules have been removed. Bot now directly binds to a
|
||||
Pipeline or Workflow. This method is kept for backward compatibility
|
||||
but only returns the direct binding.
|
||||
|
||||
Returns:
|
||||
tuple: (pipeline_uuid, routed_by_rule) - routed_by_rule is True
|
||||
when a routing rule matched, False when falling back to default.
|
||||
tuple: (pipeline_uuid, routed_by_rule) - routed_by_rule is always False
|
||||
as routing rules are no longer used.
|
||||
"""
|
||||
rules = self.bot_entity.pipeline_routing_rules or []
|
||||
element_type_set = set(message_element_types or [])
|
||||
binding_type, binding_uuid = self.get_binding_info()
|
||||
|
||||
for rule in rules:
|
||||
rule_type = rule.get('type')
|
||||
operator = rule.get('operator', 'eq')
|
||||
rule_value = rule.get('value', '')
|
||||
target_uuid = rule.get('pipeline_uuid')
|
||||
if not rule_type or not target_uuid:
|
||||
continue
|
||||
# If bound to workflow, return None for pipeline_uuid
|
||||
# The caller should check binding_type and handle accordingly
|
||||
if binding_type == 'workflow':
|
||||
# For workflow binding, we still need to return something
|
||||
# The actual workflow handling should be done by the caller
|
||||
return None, False
|
||||
|
||||
if rule_type == 'launcher_type':
|
||||
if self._match_operator(launcher_type, operator, rule_value):
|
||||
return target_uuid, True
|
||||
elif rule_type == 'launcher_id':
|
||||
if self._match_operator(str(launcher_id), operator, str(rule_value)):
|
||||
return target_uuid, True
|
||||
elif rule_type == 'message_content':
|
||||
if self._match_operator(message_text, operator, rule_value):
|
||||
return target_uuid, True
|
||||
elif rule_type == 'message_has_element':
|
||||
has_element = rule_value in element_type_set
|
||||
if operator == 'eq' and has_element:
|
||||
return target_uuid, True
|
||||
elif operator == 'neq' and not has_element:
|
||||
return target_uuid, True
|
||||
return binding_uuid, False
|
||||
|
||||
return self.bot_entity.use_pipeline_uuid, False
|
||||
async def _handle_workflow_message(
|
||||
self,
|
||||
event: platform_events.MessageEvent,
|
||||
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
|
||||
workflow_uuid: str,
|
||||
launcher_type: str,
|
||||
launcher_id: str | int,
|
||||
sender_id: str | int,
|
||||
) -> None:
|
||||
"""Handle message by executing the bound workflow directly."""
|
||||
message_content = str(event.message_chain)
|
||||
message_chain_obj = event.message_chain
|
||||
|
||||
# Build message context
|
||||
sender_name = None
|
||||
if hasattr(event, 'sender'):
|
||||
sender = event.sender
|
||||
if hasattr(sender, 'nickname'):
|
||||
sender_name = sender.nickname
|
||||
elif hasattr(sender, 'member_name'):
|
||||
sender_name = sender.member_name
|
||||
|
||||
is_group = launcher_type == 'group'
|
||||
message_context = {
|
||||
'message_id': str(getattr(event, 'message_id', '')),
|
||||
'message_content': message_content,
|
||||
'sender_id': str(sender_id),
|
||||
'sender_name': sender_name or 'User',
|
||||
'platform': adapter.__class__.__name__,
|
||||
'conversation_id': str(launcher_id),
|
||||
'is_group': is_group,
|
||||
'group_id': str(launcher_id) if is_group else None,
|
||||
'mentions': [],
|
||||
'reply_to': None,
|
||||
'raw_message': {
|
||||
'message': message_chain_obj.model_dump() if hasattr(message_chain_obj, 'model_dump') else str(message_chain_obj),
|
||||
'launcher_id': launcher_id,
|
||||
'session_type': launcher_type,
|
||||
},
|
||||
}
|
||||
|
||||
trigger_data = {
|
||||
'message': message_content,
|
||||
'message_chain': message_chain_obj.model_dump() if hasattr(message_chain_obj, 'model_dump') else str(message_chain_obj),
|
||||
'session_type': launcher_type,
|
||||
'connection_id': str(launcher_id),
|
||||
'message_context': message_context,
|
||||
}
|
||||
|
||||
session_id = f'{launcher_type}_{launcher_id}'
|
||||
logger.info(f'Processing workflow message from {session_id}: {message_content}')
|
||||
|
||||
try:
|
||||
from ..api.http.service.workflow import WorkflowExecutionFailedError
|
||||
|
||||
execution_id = await self.ap.workflow_service.execute_workflow(
|
||||
workflow_uuid=workflow_uuid,
|
||||
trigger_type='message',
|
||||
trigger_data=trigger_data,
|
||||
session_id=session_id,
|
||||
user_id=str(sender_id),
|
||||
bot_id=self.bot_entity.uuid,
|
||||
)
|
||||
except WorkflowExecutionFailedError as e:
|
||||
await self.logger.error(f'Workflow execution failed: {e.message}')
|
||||
except Exception as e:
|
||||
await self.logger.error(f'Workflow execution error: {e}')
|
||||
|
||||
async def _record_discarded_message(
|
||||
self,
|
||||
@@ -229,6 +264,20 @@ class RuntimeBot:
|
||||
|
||||
message_text = str(event.message_chain)
|
||||
element_types = [comp.type for comp in event.message_chain]
|
||||
binding_type, binding_uuid = self.get_binding_info()
|
||||
|
||||
# Handle workflow binding separately from pipeline
|
||||
if binding_type == 'workflow':
|
||||
await self._handle_workflow_message(
|
||||
event=event,
|
||||
adapter=adapter,
|
||||
workflow_uuid=binding_uuid,
|
||||
launcher_type='person',
|
||||
launcher_id=launcher_id,
|
||||
sender_id=event.sender.id,
|
||||
)
|
||||
return
|
||||
|
||||
pipeline_uuid, routed_by_rule = self.resolve_pipeline_uuid(
|
||||
'person', launcher_id, message_text, element_types
|
||||
)
|
||||
@@ -290,6 +339,20 @@ class RuntimeBot:
|
||||
|
||||
message_text = str(event.message_chain)
|
||||
element_types = [comp.type for comp in event.message_chain]
|
||||
binding_type, binding_uuid = self.get_binding_info()
|
||||
|
||||
# Handle workflow binding separately from pipeline
|
||||
if binding_type == 'workflow':
|
||||
await self._handle_workflow_message(
|
||||
event=event,
|
||||
adapter=adapter,
|
||||
workflow_uuid=binding_uuid,
|
||||
launcher_type='group',
|
||||
launcher_id=launcher_id,
|
||||
sender_id=event.sender.id,
|
||||
)
|
||||
return
|
||||
|
||||
pipeline_uuid, routed_by_rule = self.resolve_pipeline_uuid(
|
||||
'group', launcher_id, message_text, element_types
|
||||
)
|
||||
@@ -501,8 +564,6 @@ class PlatformManager:
|
||||
bot_entity.adapter_config,
|
||||
logger,
|
||||
)
|
||||
if hasattr(adapter_inst, 'ap'):
|
||||
adapter_inst.ap = self.ap
|
||||
|
||||
# 如果 adapter 支持 set_bot_uuid 方法,设置 bot_uuid(用于统一 webhook)
|
||||
if hasattr(adapter_inst, 'set_bot_uuid'):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -103,41 +103,6 @@ spec:
|
||||
type: string
|
||||
required: true
|
||||
default: "填写你的卡片template_id"
|
||||
- name: human_input_card_template_download
|
||||
label:
|
||||
en_US: Download Human Input Card Template
|
||||
zh_Hans: 下载人工输入卡片模板
|
||||
zh_Hant: 下載人工輸入卡片範本
|
||||
description:
|
||||
en_US: "Used as the only card template ID for the whole conversation turn. Download the built-in template, then import the JSON in DingTalk Open Platform > Card Platform / Card Template Management. After DingTalk creates the template, copy its template ID into the field below. The template already wires `content` (MarkdownBlock) and `btns` (ButtonGroup). Leave empty to fall back to the legacy two-card behavior."
|
||||
zh_Hans: "用作整个对话回合唯一卡片的模板 ID。先下载内置模板,再到钉钉开放平台 > 卡片平台 / 卡片模板管理中导入该 JSON;钉钉生成模板后,将模板 ID 填到这里。模板已预先连好 `content` (MarkdownBlock) 与 `btns` (ButtonGroup)。留空则降级为旧的双卡行为。"
|
||||
zh_Hant: "用作整個對話回合唯一卡片的範本 ID。先下載內建範本,再到釘釘開放平台 > 卡片平台 / 卡片範本管理中匯入該 JSON;釘釘產生範本後,將範本 ID 填到這裡。範本已預先連好 `content` (MarkdownBlock) 與 `btns` (ButtonGroup)。留空則降級為舊的雙卡行為。"
|
||||
type: download-link
|
||||
required: false
|
||||
default: ""
|
||||
url: /api/v1/platform/adapters/dingtalk/human-input-card-template
|
||||
download_filename: dingtalk_human_input_card.json
|
||||
help_links:
|
||||
zh: https://open-dev.dingtalk.com/fe/card
|
||||
en: https://open-dev.dingtalk.com/fe/card
|
||||
ja: https://open-dev.dingtalk.com/fe/card
|
||||
help_label:
|
||||
en_US: Import Guide
|
||||
zh_Hans: 导入指引
|
||||
zh_Hant: 匯入指引
|
||||
ja_JP: インポート手順
|
||||
- name: human_input_card_template_id
|
||||
label:
|
||||
en_US: Human Input Card Template ID
|
||||
zh_Hans: 人工输入卡片模板ID
|
||||
zh_Hant: 人工輸入卡片範本ID
|
||||
description:
|
||||
en_US: "Paste the template ID generated after importing the human input card template."
|
||||
zh_Hans: "填写导入人工输入卡片模板后生成的模板 ID。"
|
||||
zh_Hant: "填寫匯入人工輸入卡片範本後產生的範本 ID。"
|
||||
type: string
|
||||
required: false
|
||||
default: ""
|
||||
execution:
|
||||
python:
|
||||
path: ./dingtalk.py
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import discord
|
||||
from discord import ui as discord_ui
|
||||
|
||||
import typing
|
||||
import re
|
||||
@@ -9,8 +8,6 @@ import base64
|
||||
import uuid
|
||||
import os
|
||||
import datetime
|
||||
import time
|
||||
import traceback
|
||||
|
||||
# 使用BytesIO创建文件对象,避免路径问题
|
||||
import io
|
||||
@@ -827,69 +824,6 @@ class DiscordEventConverter(abstract_platform_adapter.AbstractEventConverter):
|
||||
)
|
||||
|
||||
|
||||
class DiscordFormView(discord_ui.View):
|
||||
"""Discord ``ui.View`` that renders one button per Dify form action.
|
||||
|
||||
Each button's click triggers ``adapter._on_form_button_click`` which
|
||||
acks the interaction, locks the buttons in place, and enqueues a
|
||||
synthetic ``_dify_form_action`` query so the runner resumes the
|
||||
workflow.
|
||||
"""
|
||||
|
||||
# Discord button style mapping for Dify ``button_style`` values.
|
||||
_STYLE_MAP: typing.ClassVar[dict] = {
|
||||
'primary': discord.ButtonStyle.primary,
|
||||
'danger': discord.ButtonStyle.danger,
|
||||
'warning': discord.ButtonStyle.danger,
|
||||
'success': discord.ButtonStyle.success,
|
||||
'default': discord.ButtonStyle.secondary,
|
||||
'': discord.ButtonStyle.secondary,
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'DiscordAdapter',
|
||||
session_key: str,
|
||||
actions: list,
|
||||
timeout: float = 1800,
|
||||
):
|
||||
super().__init__(timeout=timeout)
|
||||
self._adapter = adapter
|
||||
self._session_key = session_key
|
||||
# Discord caps a view at 25 children (5 rows × 5 buttons). Trim
|
||||
# silently — most Dify forms have ≤10 actions in practice.
|
||||
for idx, action in enumerate(actions[:25]):
|
||||
action_id = str(action.get('id') or '')
|
||||
label = str(action.get('title') or action_id or f'Option {idx + 1}')
|
||||
style = self._STYLE_MAP.get(
|
||||
str(action.get('button_style') or '').lower(),
|
||||
discord.ButtonStyle.secondary,
|
||||
)
|
||||
# custom_id must be unique within the view and ≤100 chars.
|
||||
# Encode (session, idx) so we can recover the action even
|
||||
# if Dify ids contain unsafe characters.
|
||||
custom_id = f'lb_form:{idx}:{action_id[:80]}'[:100]
|
||||
button = discord_ui.Button(
|
||||
label=label[:80], # Discord label limit
|
||||
style=style,
|
||||
custom_id=custom_id,
|
||||
)
|
||||
button.callback = self._make_callback(action_id, label)
|
||||
self.add_item(button)
|
||||
|
||||
def _make_callback(self, action_id: str, action_title: str):
|
||||
async def _cb(interaction: discord.Interaction):
|
||||
await self._adapter._on_form_button_click(
|
||||
interaction=interaction,
|
||||
session_key=self._session_key,
|
||||
action_id=action_id,
|
||||
action_title=action_title,
|
||||
view=self,
|
||||
)
|
||||
|
||||
return _cb
|
||||
|
||||
|
||||
class DiscordAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
bot: discord.Client = pydantic.Field(exclude=True)
|
||||
|
||||
@@ -903,10 +837,6 @@ class DiscordAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
|
||||
voice_manager: VoiceConnectionManager | None = pydantic.Field(exclude=True, default=None)
|
||||
|
||||
# Injected by botmgr at construction so the form-button callback can
|
||||
# enqueue a synthetic resume query (`_dify_form_action`) on the pool.
|
||||
ap: typing.Any = pydantic.Field(exclude=True, default=None)
|
||||
|
||||
def __init__(self, config: dict, logger: abstract_platform_logger.AbstractEventLogger, **kwargs):
|
||||
bot_account_id = config['client_id']
|
||||
|
||||
@@ -930,18 +860,8 @@ class DiscordAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
|
||||
args = {}
|
||||
|
||||
# Proxy: config > env var > auto-detect.
|
||||
# discord.py uses aiohttp which does NOT respect http_proxy env
|
||||
# vars by default — we must pass proxy= explicitly.
|
||||
proxy = (
|
||||
config.get('proxy')
|
||||
or os.getenv('http_proxy')
|
||||
or os.getenv('HTTP_PROXY')
|
||||
or os.getenv('https_proxy')
|
||||
or os.getenv('HTTPS_PROXY')
|
||||
)
|
||||
if proxy:
|
||||
args['proxy'] = proxy
|
||||
if os.getenv('http_proxy'):
|
||||
args['proxy'] = os.getenv('http_proxy')
|
||||
|
||||
bot = MyClient(intents=intents, **args)
|
||||
|
||||
@@ -955,19 +875,6 @@ class DiscordAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Per-resp-message-id buffer for the accumulated text yielded by
|
||||
# the runner. Discord's edit-message ratelimit (5/5s) makes true
|
||||
# progressive streaming impractical, so we collect chunks and
|
||||
# render once on is_final. ``_form_data`` on the final chunk
|
||||
# diverts to the button-view path.
|
||||
self._stream_buffer: dict[str, str] = {}
|
||||
# session_key -> {form_data, channel_id, thread_id, sender_id,
|
||||
# posted_at, view_message_id}
|
||||
# Populated when we send a form view; consumed when the user
|
||||
# clicks a button so we know which workflow_run / form_token to
|
||||
# resume.
|
||||
self._pending_forms: dict[str, dict] = {}
|
||||
|
||||
# Voice functionality methods
|
||||
async def join_voice_channel(self, guild_id: int, channel_id: int, user_id: int = None) -> discord.VoiceClient:
|
||||
"""
|
||||
@@ -1161,12 +1068,7 @@ class DiscordAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
):
|
||||
msg_to_send, files = await self.message_converter.yiri2target(message)
|
||||
|
||||
# Synthetic events (button-click resume) have no inbound discord
|
||||
# Message. Route via the channel we cached when the user clicked.
|
||||
source = message_source.source_platform_object
|
||||
if not isinstance(source, discord.Message):
|
||||
await self._reply_synthetic(message_source, msg_to_send, files)
|
||||
return
|
||||
assert isinstance(message_source.source_platform_object, discord.Message)
|
||||
|
||||
args = {
|
||||
'content': msg_to_send,
|
||||
@@ -1176,7 +1078,7 @@ class DiscordAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
args['files'] = files
|
||||
|
||||
if quote_origin:
|
||||
args['reference'] = source
|
||||
args['reference'] = message_source.source_platform_object
|
||||
|
||||
has_at = False
|
||||
|
||||
@@ -1188,422 +1090,7 @@ class DiscordAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
if has_at:
|
||||
args['mention_author'] = True
|
||||
|
||||
await source.channel.send(**args)
|
||||
|
||||
async def _reply_synthetic(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
msg_to_send: str,
|
||||
files: list,
|
||||
) -> None:
|
||||
"""Deliver a reply for a button-click-resumed (synthetic) event.
|
||||
|
||||
We don't have an inbound discord.Message to anchor to; instead
|
||||
look up the channel cached in ``_pending_forms[session_key +
|
||||
'__last_channel']`` from the most recent button click.
|
||||
"""
|
||||
if isinstance(message_source, platform_events.GroupMessage):
|
||||
# _handle_form_chunk uses channel_id alone as the session
|
||||
# scope, and launcher_id was set to channel_id when
|
||||
# synthesizing the event.
|
||||
session_key = f'c:{message_source.group.id}'
|
||||
else:
|
||||
session_key = f'p:{message_source.sender.id}'
|
||||
|
||||
cached = self._pending_forms.get(session_key + '__last_channel') or {}
|
||||
channel = cached.get('channel')
|
||||
if channel is None:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.warning(
|
||||
f'Discord: synthetic reply has no cached channel for '
|
||||
f'{session_key}; dropping content (len={len(msg_to_send)})'
|
||||
)
|
||||
return
|
||||
|
||||
args: dict[str, typing.Any] = {'content': msg_to_send}
|
||||
if files:
|
||||
args['files'] = files
|
||||
try:
|
||||
await channel.send(**args)
|
||||
except Exception:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.error(f'Discord: synthetic reply send failed: {traceback.format_exc()}')
|
||||
|
||||
# Discord allows 5 edits per 5 seconds per message. We throttle
|
||||
# to one edit per 8 runner-chunks (runner already yields every 8
|
||||
# text_chunks internally), which stays comfortably within limits.
|
||||
_STREAM_EDIT_INTERVAL = 8
|
||||
|
||||
async def is_stream_output_supported(self) -> bool:
|
||||
return True
|
||||
|
||||
async def create_message_card(self, message_id: str, event: platform_events.MessageEvent) -> bool:
|
||||
"""Set up a stream context for progressive editing.
|
||||
|
||||
The first non-empty reply_message_chunk will send the initial
|
||||
message; subsequent chunks edit it in place.
|
||||
"""
|
||||
source = event.source_platform_object
|
||||
if not isinstance(source, discord.Message):
|
||||
return False
|
||||
self._stream_buffer[message_id] = {
|
||||
'channel': source.channel,
|
||||
'sent_message': None, # discord.Message set on first send
|
||||
'last_content': '',
|
||||
'chunk_count': 0,
|
||||
}
|
||||
return True
|
||||
|
||||
async def reply_message_chunk(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
bot_message: typing.Any,
|
||||
message: platform_message.MessageChain,
|
||||
quote_origin: bool = False,
|
||||
is_final: bool = False,
|
||||
):
|
||||
msg_id = (
|
||||
bot_message.get('resp_message_id')
|
||||
if isinstance(bot_message, dict)
|
||||
else getattr(bot_message, 'resp_message_id', None)
|
||||
)
|
||||
|
||||
text_parts = [m.text for m in message if isinstance(m, platform_message.Plain)]
|
||||
chunk_text = '\n\n'.join(t for t in text_parts if t)
|
||||
|
||||
form_data = getattr(bot_message, '_form_data', None) if not isinstance(bot_message, dict) else None
|
||||
|
||||
ctx = self._stream_buffer.get(msg_id) if msg_id else None
|
||||
|
||||
# If the stream ctx was not set up (create_message_card wasn't
|
||||
# called, e.g. synthetic event), or the final chunk carries a
|
||||
# form, skip progressive editing entirely.
|
||||
if ctx is None or form_data:
|
||||
try:
|
||||
if form_data and is_final:
|
||||
await self._handle_form_chunk(message_source, form_data)
|
||||
elif is_final and chunk_text:
|
||||
await self.reply_message(
|
||||
message_source,
|
||||
platform_message.MessageChain([platform_message.Plain(text=chunk_text)]),
|
||||
quote_origin,
|
||||
)
|
||||
finally:
|
||||
self._stream_buffer.pop(msg_id, None)
|
||||
return
|
||||
|
||||
# Progressive streaming path: send first chunk, edit subsequent.
|
||||
ctx['chunk_count'] += 1
|
||||
|
||||
# Runner yields the full accumulated text on each chunk, so we
|
||||
# always replace (not append).
|
||||
if chunk_text:
|
||||
ctx['last_content'] = chunk_text
|
||||
|
||||
sent = ctx['sent_message']
|
||||
|
||||
if sent is None:
|
||||
# First non-empty chunk — send the initial message.
|
||||
if not ctx['last_content']:
|
||||
return # No content yet, wait for next chunk.
|
||||
try:
|
||||
sent = await ctx['channel'].send(ctx['last_content'])
|
||||
ctx['sent_message'] = sent
|
||||
except Exception:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.error(f'Discord stream send failed: {traceback.format_exc()}')
|
||||
self._stream_buffer.pop(msg_id, None)
|
||||
return
|
||||
|
||||
if is_final:
|
||||
# Final chunk — edit to the full content, then clean up.
|
||||
if ctx['last_content'] and ctx['last_content'] != sent.content:
|
||||
try:
|
||||
await sent.edit(content=ctx['last_content'][:2000])
|
||||
except Exception:
|
||||
pass # Best-effort
|
||||
self._stream_buffer.pop(msg_id, None)
|
||||
elif (ctx['chunk_count'] % self._STREAM_EDIT_INTERVAL) == 0:
|
||||
# Intermediate edit — throttle to avoid rate limits.
|
||||
if ctx['last_content'] and ctx['last_content'] != sent.content:
|
||||
try:
|
||||
await sent.edit(content=ctx['last_content'][:2000])
|
||||
except Exception:
|
||||
pass # Rate-limited or deleted — ignore.
|
||||
|
||||
async def _handle_form_chunk(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
form_data: dict,
|
||||
) -> None:
|
||||
"""Render a Dify form pause as a Discord embed + button View.
|
||||
|
||||
Mirrors the QQ / Telegram / Lark form path: the button's click
|
||||
callback synthesizes a ``_dify_form_action`` query so the runner's
|
||||
``_merge_pending_form_action`` resumes the workflow.
|
||||
"""
|
||||
source = message_source.source_platform_object
|
||||
|
||||
actions = form_data.get('actions') or []
|
||||
if not actions:
|
||||
# Nothing clickable — fall back to plain text.
|
||||
if source is not None:
|
||||
await self.reply_message(
|
||||
message_source,
|
||||
platform_message.MessageChain(
|
||||
[platform_message.Plain(text=str(form_data.get('node_title') or ''))]
|
||||
),
|
||||
)
|
||||
return
|
||||
|
||||
node_title = str(form_data.get('node_title') or 'Confirmation needed')
|
||||
form_content = str(form_data.get('form_content') or '').strip()
|
||||
|
||||
# Two paths:
|
||||
# (a) Real message — extract channel from source.
|
||||
# (b) Synthetic event (button-click resume) — no
|
||||
# source_platform_object; recover the channel we cached
|
||||
# when the user clicked.
|
||||
if isinstance(source, discord.Message):
|
||||
channel = source.channel
|
||||
guild_id = str(source.guild.id) if source.guild else ''
|
||||
sender_id = str(source.author.id)
|
||||
channel_id = str(source.channel.id)
|
||||
session_key = f'c:{channel_id}' if guild_id else f'p:{sender_id}'
|
||||
else:
|
||||
# Synthetic event — resolve session_key from event shape,
|
||||
# then look up the cached channel from the click.
|
||||
if isinstance(message_source, platform_events.GroupMessage):
|
||||
# launcher_id was set to channel_id when we synthesized.
|
||||
channel_id = str(message_source.group.id)
|
||||
session_key = f'c:{channel_id}'
|
||||
else:
|
||||
session_key = f'p:{message_source.sender.id}'
|
||||
channel_id = ''
|
||||
|
||||
cached = self._pending_forms.get(session_key + '__last_channel')
|
||||
channel = cached.get('channel') if cached else None
|
||||
guild_id = (cached or {}).get('guild_id', '')
|
||||
sender_id = str(message_source.sender.id) if message_source.sender else ''
|
||||
if channel is None:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.warning(
|
||||
f'Discord: synthetic form chunk has no cached channel for '
|
||||
f'{session_key}; cannot render form buttons'
|
||||
)
|
||||
return
|
||||
|
||||
body_parts: list[str] = []
|
||||
if form_content:
|
||||
body_parts.append(form_content)
|
||||
embed_body = '\n\n'.join(body_parts)
|
||||
# Discord embed.description has a 4096 char limit — defensive trim.
|
||||
if len(embed_body) > 4000:
|
||||
embed_body = embed_body[:3990] + '\n\n…(truncated)'
|
||||
|
||||
embed = discord.Embed(
|
||||
title=node_title[:256],
|
||||
description=embed_body,
|
||||
color=discord.Color.blurple(),
|
||||
)
|
||||
|
||||
view = DiscordFormView(
|
||||
adapter=self,
|
||||
session_key=session_key,
|
||||
actions=actions,
|
||||
timeout=1800, # 30 min — matches Dify form_token TTL
|
||||
)
|
||||
|
||||
try:
|
||||
sent_msg = await channel.send(embed=embed, view=view)
|
||||
except Exception:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.error(f'Discord: form view send failed: {traceback.format_exc()}')
|
||||
return
|
||||
|
||||
self._pending_forms[session_key] = {
|
||||
'form_data': form_data,
|
||||
'channel_id': channel_id,
|
||||
'guild_id': guild_id,
|
||||
'sender_id': sender_id,
|
||||
'view_message_id': str(sent_msg.id),
|
||||
'posted_at': time.time(),
|
||||
}
|
||||
|
||||
if self.ap is not None:
|
||||
self.ap.logger.info(f'Discord: form view posted session={session_key} actions={len(actions)}')
|
||||
|
||||
async def _on_form_button_click(
|
||||
self,
|
||||
interaction: discord.Interaction,
|
||||
session_key: str,
|
||||
action_id: str,
|
||||
action_title: str,
|
||||
view: DiscordFormView,
|
||||
) -> None:
|
||||
"""Handle a click on a form button — ack, resume the workflow,
|
||||
and disable the View buttons so the choice is visually locked in."""
|
||||
import langbot_plugin.api.entities.builtin.provider.session as provider_session
|
||||
|
||||
# ACK first (3-second deadline before Discord shows "interaction failed").
|
||||
try:
|
||||
await interaction.response.defer()
|
||||
except discord.HTTPException:
|
||||
# Already responded somehow — proceed regardless.
|
||||
pass
|
||||
|
||||
pending = self._pending_forms.get(session_key)
|
||||
if not pending:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.warning(
|
||||
f'Discord: button click on stale session {session_key}; ignoring (action_id={action_id!r})'
|
||||
)
|
||||
await self._lock_view_message(interaction, view, action_title, stale=True)
|
||||
return
|
||||
|
||||
form_data: dict = pending.get('form_data') or {}
|
||||
guild_id = pending.get('guild_id', '')
|
||||
channel_id = pending.get('channel_id', '')
|
||||
initiator_id = str(pending.get('sender_id', '') or '')
|
||||
actor_id = str(interaction.user.id) if interaction.user is not None else initiator_id
|
||||
if not guild_id and initiator_id and actor_id != initiator_id:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.warning(
|
||||
f'Discord: user {actor_id} cannot act on private form created for {initiator_id}'
|
||||
)
|
||||
await self._lock_view_message(interaction, view, action_title, stale=True)
|
||||
return
|
||||
|
||||
self._pending_forms.pop(session_key, None)
|
||||
|
||||
# Lock the buttons in place: disable everything, mark chosen one.
|
||||
await self._lock_view_message(interaction, view, action_title)
|
||||
|
||||
# In group context the launcher remains the channel so Dify resumes
|
||||
# the original group session. The synthetic sender is still the real
|
||||
# clicker, preserving actor identity for auditing and routing rules.
|
||||
if guild_id:
|
||||
launcher_type = provider_session.LauncherTypes.GROUP
|
||||
launcher_id = channel_id
|
||||
else:
|
||||
launcher_type = provider_session.LauncherTypes.PERSON
|
||||
launcher_id = initiator_id or actor_id
|
||||
|
||||
form_action_data = {
|
||||
'form_token': form_data.get('form_token', ''),
|
||||
'workflow_run_id': form_data.get('workflow_run_id', ''),
|
||||
'action_id': action_id,
|
||||
'action_title': action_title,
|
||||
'node_title': form_data.get('node_title', ''),
|
||||
'user': f'{launcher_type.value}_{launcher_id}',
|
||||
'inputs': {},
|
||||
}
|
||||
|
||||
message_chain = platform_message.MessageChain([platform_message.Plain(text=f'[Form Action: {action_title}]')])
|
||||
|
||||
# Synthesize a platform event so the pipeline can run the resume
|
||||
# query. source_platform_object=None signals "no inbound discord
|
||||
# message" — reply_message must tolerate this (it falls through
|
||||
# to channel.send via the cached interaction.channel below).
|
||||
if launcher_type == provider_session.LauncherTypes.GROUP:
|
||||
synthetic_event: platform_events.MessageEvent = platform_events.GroupMessage(
|
||||
sender=platform_entities.GroupMember(
|
||||
id=actor_id,
|
||||
member_name=interaction.user.display_name if interaction.user else '',
|
||||
permission='MEMBER',
|
||||
group=platform_entities.Group(
|
||||
id=launcher_id,
|
||||
name=channel_id,
|
||||
permission=platform_entities.Permission.Member,
|
||||
),
|
||||
special_title='',
|
||||
),
|
||||
message_chain=message_chain,
|
||||
time=int(time.time()),
|
||||
source_platform_object=None,
|
||||
)
|
||||
else:
|
||||
synthetic_event = platform_events.FriendMessage(
|
||||
sender=platform_entities.Friend(
|
||||
id=actor_id,
|
||||
nickname=interaction.user.display_name if interaction.user else '',
|
||||
remark='',
|
||||
),
|
||||
message_chain=message_chain,
|
||||
time=int(time.time()),
|
||||
source_platform_object=None,
|
||||
)
|
||||
|
||||
if self.ap is None:
|
||||
if self.logger:
|
||||
await self.logger.error('Discord: ap not injected; cannot enqueue button-click query')
|
||||
return
|
||||
|
||||
bot_uuid = ''
|
||||
pipeline_uuid = form_data.get('pipeline_uuid') or None
|
||||
for bot in self.ap.platform_mgr.bots:
|
||||
if bot.adapter is self:
|
||||
bot_uuid = bot.bot_entity.uuid
|
||||
pipeline_uuid = pipeline_uuid or bot.bot_entity.use_pipeline_uuid
|
||||
break
|
||||
|
||||
# Remember the channel so _reply_synthetic and _handle_form_chunk
|
||||
# (synthetic-event path) can find a target. guild_id is needed
|
||||
# to reconstruct the launcher_type on subsequent form pauses.
|
||||
self._pending_forms[session_key + '__last_channel'] = {
|
||||
'channel': interaction.channel,
|
||||
'guild_id': guild_id,
|
||||
'posted_at': time.time(),
|
||||
}
|
||||
|
||||
try:
|
||||
await self.ap.query_pool.add_query(
|
||||
bot_uuid=bot_uuid,
|
||||
launcher_type=launcher_type,
|
||||
launcher_id=launcher_id,
|
||||
sender_id=actor_id,
|
||||
message_event=synthetic_event,
|
||||
message_chain=message_chain,
|
||||
adapter=self,
|
||||
pipeline_uuid=pipeline_uuid,
|
||||
variables={
|
||||
'_dify_form_action': form_action_data,
|
||||
'_routed_by_rule': True,
|
||||
},
|
||||
)
|
||||
if self.ap is not None:
|
||||
self.ap.logger.info(
|
||||
f'Discord: button-click query enqueued action_id={action_id!r} '
|
||||
f'session={session_key} actor_id={actor_id}'
|
||||
)
|
||||
except Exception:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.error(f'Discord: enqueue button-click query failed: {traceback.format_exc()}')
|
||||
|
||||
async def _lock_view_message(
|
||||
self,
|
||||
interaction: discord.Interaction,
|
||||
view: DiscordFormView,
|
||||
chosen_title: str,
|
||||
stale: bool = False,
|
||||
) -> None:
|
||||
"""Disable all buttons on the form view and annotate the chosen
|
||||
one — mirrors DingTalk/Lark's in-card selection feedback."""
|
||||
try:
|
||||
for child in view.children:
|
||||
if not isinstance(child, discord_ui.Button):
|
||||
continue
|
||||
child.disabled = True
|
||||
if not stale and child.label == chosen_title:
|
||||
child.style = discord.ButtonStyle.success
|
||||
if not (child.label or '').startswith('✓ '):
|
||||
child.label = f'✓ {child.label}'
|
||||
view.stop()
|
||||
if interaction.message is not None:
|
||||
await interaction.message.edit(view=view)
|
||||
except Exception:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.warning(f'Discord: lock-view-message failed (non-fatal): {traceback.format_exc()}')
|
||||
await message_source.source_platform_object.channel.send(**args)
|
||||
|
||||
async def is_muted(self, group_id: int) -> bool:
|
||||
return False
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -11,13 +11,7 @@ import langbot_plugin.api.definition.abstract.platform.adapter as abstract_platf
|
||||
import langbot_plugin.api.entities.builtin.platform.message as platform_message
|
||||
import langbot_plugin.api.entities.builtin.platform.events as platform_events
|
||||
import langbot_plugin.api.entities.builtin.platform.entities as platform_entities
|
||||
from langbot.libs.qq_official_api.api import (
|
||||
QQ_SELECT_ACTION_PREFIX,
|
||||
QQOfficialClient,
|
||||
build_keyboard_from_form,
|
||||
build_keyboard_from_select_field,
|
||||
resolve_select_button_action,
|
||||
)
|
||||
from langbot.libs.qq_official_api.api import QQOfficialClient
|
||||
from langbot.libs.qq_official_api.qqofficialevent import QQOfficialEvent
|
||||
from ...utils import image
|
||||
from ..logger import EventLogger
|
||||
@@ -197,7 +191,6 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
enable_webhook: bool = False
|
||||
message_converter: QQOfficialMessageConverter = QQOfficialMessageConverter()
|
||||
event_converter: QQOfficialEventConverter = QQOfficialEventConverter()
|
||||
ap: typing.Any = None
|
||||
|
||||
def __init__(self, config: dict, logger: EventLogger):
|
||||
enable_webhook = config.get('enable-webhook', False)
|
||||
@@ -223,31 +216,6 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
self._stream_ctx_ts: dict[str, float] = {}
|
||||
self._fallback_text: dict[str, str] = {}
|
||||
self._fallback_text_ts: dict[str, float] = {}
|
||||
# Dify form-action bookkeeping for the human-input button flow.
|
||||
# session_key = "<scene>_<id>" where scene is c2c/group/channel and
|
||||
# id is user_openid / group_openid / channel_id.
|
||||
# session_key -> {form_data, msg_id, event_id, scene, target_id,
|
||||
# sender_id, posted_at}
|
||||
# Set when we send a markdown+keyboard card and consulted when:
|
||||
# (a) INTERACTION_CREATE fires — we look up the form by
|
||||
# session_key (button's `data` carries the action_id),
|
||||
# (b) the resumed-workflow query needs to find a passive-reply
|
||||
# event_id (INTERACTION_CREATE id, 30-min validity).
|
||||
self._pending_forms: dict[str, dict] = {}
|
||||
# session_key -> most recent ``INTERACTION_CREATE`` event_id, used
|
||||
# as the passive event_id for the resumed query's LLM output.
|
||||
self._session_event_ids: dict[str, dict] = {}
|
||||
# Per-anchor msg_seq counter. QQ accepts up to 5 passive replies
|
||||
# per (msg_id|event_id) within 60 min, but each reuse needs a
|
||||
# fresh ``msg_seq`` — re-sending with msg_seq=1 is silently dedup'd.
|
||||
self._anchor_msg_seq: dict[str, int] = {}
|
||||
|
||||
# Wire button-click handler so webhook mode catches INTERACTION_CREATE.
|
||||
# (ws mode is wired separately via on_event in _run_websocket so the
|
||||
# raw payload bypasses get_message's message-only flattening.)
|
||||
@self.bot.on_interaction()
|
||||
async def _on_interaction(event_data: dict, interaction_id: typing.Optional[str]):
|
||||
await self._handle_interaction_create(event_data, interaction_id)
|
||||
|
||||
async def reply_message(
|
||||
self,
|
||||
@@ -259,13 +227,6 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
message_source,
|
||||
)
|
||||
|
||||
# Synthetic event (button-click resume): no inbound platform
|
||||
# object → no msg_id. Route via the cached INTERACTION_CREATE
|
||||
# event_id (valid 30 min, no quota cost).
|
||||
if qq_official_event is None:
|
||||
await self._reply_synthetic(message_source, message)
|
||||
return
|
||||
|
||||
content_list = await QQOfficialMessageConverter.yiri2target(message)
|
||||
|
||||
# 确定 target_type 和 target_id
|
||||
@@ -415,9 +376,6 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
await self.logger.info('QQ Official WebSocket connected and ready')
|
||||
|
||||
async def on_event(event_type: str, event_data: dict):
|
||||
# INTERACTION_CREATE is dispatched via bot.on_interaction()
|
||||
# (registered in __init__) so we get the top-level ws_event_id
|
||||
# — needed as the passive-reply event_id. It never reaches here.
|
||||
# 只处理消息事件,忽略 READY/RESUMED 等系统事件
|
||||
message_event_types = {
|
||||
'C2C_MESSAGE_CREATE',
|
||||
@@ -479,36 +437,12 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
async def is_stream_output_supported(self) -> bool:
|
||||
return self.config.get('enable-stream-reply', False)
|
||||
|
||||
@staticmethod
|
||||
def _is_form_placeholder_chunk(text: str) -> bool:
|
||||
"""Return True for invisible placeholder chunks used to carry forms."""
|
||||
|
||||
if not text:
|
||||
return False
|
||||
|
||||
cleaned = text.replace('\u200b', '').replace('\u200c', '').replace('\u200d', '').replace('\ufeff', '').strip()
|
||||
# Some Windows consoles/logs display the zero-width placeholder as
|
||||
# mojibake. Treat those variants as the same non-user-facing marker.
|
||||
return cleaned in {'', '鈥?', '​'}
|
||||
|
||||
async def create_message_card(self, message_id: str, event: platform_events.MessageEvent) -> bool:
|
||||
source = event.source_platform_object
|
||||
# Synthetic events (button-click resume) have no source object —
|
||||
# they ride a cached INTERACTION_CREATE event_id, not a streamable
|
||||
# msg_id. Skip stream setup; reply_message handles the one-shot
|
||||
# send at is_final.
|
||||
if source is None:
|
||||
return False
|
||||
# Streaming API only supports C2C private chat
|
||||
if source.t != 'C2C_MESSAGE_CREATE':
|
||||
return False
|
||||
|
||||
# The stream endpoint still consumes msg_seq for this inbound msg_id.
|
||||
# Keep the passive-reply counter in sync so a follow-up form card uses
|
||||
# msg_seq=2 instead of being deduplicated by QQ as another seq=1 send.
|
||||
if source.d_id:
|
||||
self._anchor_msg_seq[source.d_id] = max(self._anchor_msg_seq.get(source.d_id, 0), 1)
|
||||
|
||||
ctx = {
|
||||
'user_openid': source.user_openid,
|
||||
'msg_id': source.d_id,
|
||||
@@ -535,38 +469,12 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
):
|
||||
# Periodically clean up stale stream contexts
|
||||
await self._cleanup_stale_streams()
|
||||
|
||||
# Dify human-input pause: when the runner attaches `_form_data` to
|
||||
# the final chunk, finalize any in-flight stream session and send
|
||||
# a markdown + keyboard message instead. Plain-text content from
|
||||
# earlier chunks is already on the stream; we close it cleanly
|
||||
# and the buttons land as a separate reply.
|
||||
form_data = getattr(bot_message, '_form_data', None) if not isinstance(bot_message, dict) else None
|
||||
if is_final:
|
||||
_resume = getattr(bot_message, '_resume_from_form', None) if not isinstance(bot_message, dict) else None
|
||||
_open_new = getattr(bot_message, '_open_new_card', None) if not isinstance(bot_message, dict) else None
|
||||
if self.ap is not None:
|
||||
self.ap.logger.info(
|
||||
f'QQ Official reply_message_chunk final: '
|
||||
f'type={type(bot_message).__name__} '
|
||||
f'is_final={is_final} '
|
||||
f'form_data_present={form_data is not None} '
|
||||
f'resume_from_form={_resume} open_new_card={_open_new} '
|
||||
f'content_len={len(getattr(bot_message, "content", "") or "")}'
|
||||
)
|
||||
if form_data and is_final:
|
||||
await self._handle_form_chunk(message_source, message, form_data)
|
||||
return
|
||||
|
||||
# 提取纯文本内容(当前 chunk 的文本)
|
||||
text_parts = []
|
||||
for msg in message:
|
||||
if type(msg) is platform_message.Plain:
|
||||
text_parts.append(msg.text)
|
||||
chunk_text = '\n\n'.join(text_parts)
|
||||
if self._is_form_placeholder_chunk(chunk_text):
|
||||
await self.logger.debug('QQ Official: skipped invisible form placeholder chunk')
|
||||
return
|
||||
|
||||
message_id = (
|
||||
bot_message.get('resp_message_id')
|
||||
@@ -576,8 +484,7 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
if not message_id or message_id not in self._stream_ctx:
|
||||
# 非流式场景(如群聊不支持流式),累积文本后一次性回复
|
||||
if chunk_text:
|
||||
# Chunks carry the latest full snapshot, not a text delta.
|
||||
self._fallback_text[message_id] = chunk_text
|
||||
self._fallback_text[message_id] = self._fallback_text.get(message_id, '') + chunk_text
|
||||
self._fallback_text_ts[message_id] = time.time()
|
||||
if is_final:
|
||||
full_text = self._fallback_text.pop(message_id, '')
|
||||
@@ -590,7 +497,7 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
|
||||
# 累积文本
|
||||
if chunk_text:
|
||||
ctx['accumulated_text'] = chunk_text
|
||||
ctx['accumulated_text'] += chunk_text
|
||||
|
||||
# 未启动会话时,等第一个有内容的 chunk 来建立会话
|
||||
if not ctx['session_started']:
|
||||
@@ -650,489 +557,3 @@ class QQOfficialAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter
|
||||
],
|
||||
):
|
||||
return super().unregister_listener(event_type, callback)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Dify human-input button-interaction support
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
_PENDING_FORM_TTL = 1800 # 30 min — matches QQ passive-reply window.
|
||||
_MAX_REPLIES_PER_ANCHOR = 5 # QQ hard limit per msg_id / event_id.
|
||||
|
||||
def _next_msg_seq(self, anchor: str) -> typing.Optional[int]:
|
||||
"""Return the next msg_seq for an anchor, or ``None`` if the
|
||||
anchor has already been used 5 times (further sends would be
|
||||
silently dropped by QQ)."""
|
||||
if not anchor:
|
||||
return 1
|
||||
used = self._anchor_msg_seq.get(anchor, 0)
|
||||
if used >= self._MAX_REPLIES_PER_ANCHOR:
|
||||
return None
|
||||
self._anchor_msg_seq[anchor] = used + 1
|
||||
return used + 1
|
||||
|
||||
async def _reply_synthetic(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
message: platform_message.MessageChain,
|
||||
) -> None:
|
||||
"""Deliver a reply for a synthetic (button-click-resume) event.
|
||||
|
||||
Synthetic events have ``source_platform_object=None`` and no
|
||||
fresh inbound msg_id. The previous INTERACTION_CREATE id we
|
||||
cached in :attr:`_session_event_ids` is a valid passive-reply
|
||||
anchor (``event_id``) for up to 30 minutes — use it.
|
||||
"""
|
||||
if isinstance(message_source, platform_events.GroupMessage):
|
||||
target_type = 'group'
|
||||
group = getattr(message_source, 'group', None) or (
|
||||
message_source.sender.group if hasattr(message_source.sender, 'group') else None
|
||||
)
|
||||
target_id = str(group.id) if group else None
|
||||
else:
|
||||
target_type = 'c2c'
|
||||
target_id = str(message_source.sender.id) if message_source.sender else None
|
||||
|
||||
if not target_id:
|
||||
await self.logger.warning('QQ Official: synthetic reply has no target_id; dropping')
|
||||
return
|
||||
|
||||
session_key = f'{target_type}_{target_id}'
|
||||
cached = self._session_event_ids.get(session_key)
|
||||
event_id = cached.get('event_id') if cached else None
|
||||
if cached and (time.time() - cached.get('posted_at', 0)) > self._PENDING_FORM_TTL:
|
||||
event_id = None
|
||||
|
||||
if not event_id:
|
||||
await self.logger.warning(
|
||||
f'QQ Official: no cached event_id for {session_key}; '
|
||||
f'cannot deliver synthetic reply within passive-reply window'
|
||||
)
|
||||
return
|
||||
|
||||
content_list = await QQOfficialMessageConverter.yiri2target(message)
|
||||
text_parts = [c['content'] for c in content_list if c.get('type') == 'text' and c.get('content')]
|
||||
if not text_parts:
|
||||
await self.logger.info('QQ Official: synthetic reply has no text content; skipping')
|
||||
return
|
||||
text = '\n\n'.join(text_parts)
|
||||
|
||||
msg_seq = self._next_msg_seq(event_id)
|
||||
if msg_seq is None:
|
||||
await self.logger.warning(
|
||||
f'QQ Official: anchor {event_id!r} exhausted (>5 passive replies); '
|
||||
f'cannot deliver synthetic reply for {session_key}'
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
if target_type == 'c2c':
|
||||
await self.bot.send_private_text_msg(
|
||||
user_openid=target_id,
|
||||
content=text,
|
||||
event_id=event_id,
|
||||
msg_seq=msg_seq,
|
||||
)
|
||||
elif target_type == 'group':
|
||||
await self.bot.send_group_text_msg(
|
||||
group_openid=target_id,
|
||||
content=text,
|
||||
event_id=event_id,
|
||||
msg_seq=msg_seq,
|
||||
)
|
||||
except Exception:
|
||||
await self.logger.error(f'QQ Official: synthetic reply delivery failed: {traceback.format_exc()}')
|
||||
|
||||
def _resolve_target_from_source(self, source: QQOfficialEvent) -> typing.Optional[tuple[str, str]]:
|
||||
"""Return ``(target_type, target_id)`` for sending a reply, or
|
||||
``None`` if the scene cannot host a markdown+keyboard message."""
|
||||
if source is None:
|
||||
return None
|
||||
if source.t == 'C2C_MESSAGE_CREATE':
|
||||
return 'c2c', source.user_openid
|
||||
if source.t == 'GROUP_AT_MESSAGE_CREATE':
|
||||
return 'group', source.group_openid
|
||||
if source.t == 'AT_MESSAGE_CREATE':
|
||||
return 'channel', source.channel_id
|
||||
# DIRECT_MESSAGE_CREATE uses the guild DM API which does not accept
|
||||
# markdown+keyboard at the time of writing — caller falls back to text.
|
||||
return None
|
||||
|
||||
def _resolve_target_from_event(
|
||||
self, message_source: platform_events.MessageEvent
|
||||
) -> typing.Optional[tuple[str, str]]:
|
||||
"""Resolve ``(target_type, target_id)`` from the public event.
|
||||
|
||||
Prefers the platform-native source when present; falls back to
|
||||
the synthesized event's sender/group fields so button-click
|
||||
resume queries can still find a destination.
|
||||
"""
|
||||
source = message_source.source_platform_object
|
||||
if source is not None:
|
||||
return self._resolve_target_from_source(source)
|
||||
if isinstance(message_source, platform_events.GroupMessage):
|
||||
group = getattr(message_source, 'group', None) or (
|
||||
message_source.sender.group
|
||||
if message_source.sender and hasattr(message_source.sender, 'group')
|
||||
else None
|
||||
)
|
||||
if group and getattr(group, 'id', None):
|
||||
return 'group', str(group.id)
|
||||
if isinstance(message_source, platform_events.FriendMessage):
|
||||
if message_source.sender and getattr(message_source.sender, 'id', None):
|
||||
return 'c2c', str(message_source.sender.id)
|
||||
return None
|
||||
|
||||
def _prune_pending_forms(self) -> None:
|
||||
now = time.time()
|
||||
stale = [k for k, v in self._pending_forms.items() if now - v.get('posted_at', 0) > self._PENDING_FORM_TTL]
|
||||
for k in stale:
|
||||
self._pending_forms.pop(k, None)
|
||||
stale_e = [
|
||||
k for k, v in self._session_event_ids.items() if now - v.get('posted_at', 0) > self._PENDING_FORM_TTL
|
||||
]
|
||||
for k in stale_e:
|
||||
self._session_event_ids.pop(k, None)
|
||||
|
||||
async def _handle_form_chunk(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
message: platform_message.MessageChain,
|
||||
form_data: dict,
|
||||
) -> None:
|
||||
"""Send the markdown + keyboard form prompt for a Dify pause.
|
||||
|
||||
Called from ``reply_message_chunk`` when the runner attaches
|
||||
``_form_data`` to the final chunk. Replaces what would otherwise
|
||||
be a plain-text numbered-list fallback.
|
||||
"""
|
||||
if self.ap is not None:
|
||||
self.ap.logger.info(
|
||||
f'QQ Official _handle_form_chunk entered; '
|
||||
f'source_present={message_source.source_platform_object is not None} '
|
||||
f'form_actions={len(form_data.get("actions") or [])}'
|
||||
)
|
||||
self._prune_pending_forms()
|
||||
|
||||
source = message_source.source_platform_object
|
||||
scene_target = self._resolve_target_from_event(message_source)
|
||||
if scene_target is None:
|
||||
# No rich-UI fit — fall through to existing text path.
|
||||
await self.logger.info('QQ Official: form chunk on unsupported scene; falling back to text')
|
||||
text_parts = [m.text for m in message if type(m) is platform_message.Plain]
|
||||
fallback_msg = platform_message.MessageChain([platform_message.Plain(text='\n\n'.join(text_parts))])
|
||||
try:
|
||||
await self.reply_message(message_source, fallback_msg)
|
||||
except Exception:
|
||||
await self.logger.error(f'QQ Official: form fallback text send failed: {traceback.format_exc()}')
|
||||
return
|
||||
|
||||
target_type, target_id = scene_target
|
||||
session_key = f'{target_type}_{target_id}'
|
||||
|
||||
# Cancel any in-flight stream / fallback ctx so plain-text prefix
|
||||
# doesn't continue alongside the keyboard message.
|
||||
msg_id = getattr(source, 'd_id', '') or '' if source is not None else ''
|
||||
if msg_id:
|
||||
self._stream_ctx.pop(msg_id, None)
|
||||
self._stream_ctx_ts.pop(msg_id, None)
|
||||
self._fallback_text.pop(msg_id, None)
|
||||
self._fallback_text_ts.pop(msg_id, None)
|
||||
|
||||
node_title = form_data.get('node_title') or 'Confirmation needed'
|
||||
form_content = form_data.get('form_content') or ''
|
||||
is_field_step = bool(form_data.get('_current_input_field')) and not form_data.get('_action_select_only')
|
||||
parts = [f'### {node_title}']
|
||||
plain_parts = [node_title]
|
||||
if form_content.strip():
|
||||
parts.append(form_content.strip())
|
||||
plain_parts.append(form_content.strip())
|
||||
markdown_content = '\n\n'.join(parts)
|
||||
plain_content = '\n\n'.join(plain_parts)
|
||||
|
||||
keyboard = build_keyboard_from_select_field(form_data) if is_field_step else None
|
||||
is_text_field_step = is_field_step and not keyboard.get('content', {}).get('rows')
|
||||
if is_text_field_step:
|
||||
keyboard = None
|
||||
if keyboard is None and not is_text_field_step:
|
||||
keyboard = build_keyboard_from_form(form_data, buttons_per_row=2)
|
||||
if keyboard is not None and not keyboard.get('content', {}).get('rows') and not is_text_field_step:
|
||||
# No actions to render — fall back to plain text.
|
||||
text_msg = platform_message.MessageChain([platform_message.Plain(text=plain_content)])
|
||||
try:
|
||||
await self.reply_message(message_source, text_msg)
|
||||
except Exception:
|
||||
await self.logger.error(f'QQ Official: empty-keyboard fallback send failed: {traceback.format_exc()}')
|
||||
return
|
||||
|
||||
# Prefer the inbound msg_id (no quota cost). If the source is a
|
||||
# synthetic event from a prior click, the cached interaction id
|
||||
# serves as event_id for up to 30 min.
|
||||
event_id = None
|
||||
if not msg_id:
|
||||
cached = self._session_event_ids.get(session_key)
|
||||
if cached and (time.time() - cached.get('posted_at', 0)) < self._PENDING_FORM_TTL:
|
||||
event_id = cached.get('event_id')
|
||||
|
||||
anchor = msg_id or event_id or ''
|
||||
msg_seq = self._next_msg_seq(anchor)
|
||||
if msg_seq is None:
|
||||
await self.logger.warning(
|
||||
f'QQ Official: anchor {anchor!r} exhausted (>5 passive replies); '
|
||||
f'cannot deliver form card for session={session_key}'
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
await self.bot.send_markdown_keyboard(
|
||||
target_type=target_type,
|
||||
target_id=target_id,
|
||||
markdown_content=markdown_content,
|
||||
keyboard=keyboard,
|
||||
msg_id=msg_id if (msg_id and not event_id) else None,
|
||||
event_id=event_id,
|
||||
msg_seq=msg_seq,
|
||||
)
|
||||
if self.ap is not None:
|
||||
self.ap.logger.info(
|
||||
f'QQ Official: form card sent '
|
||||
f'target={target_type}/{target_id} '
|
||||
f'msg_id={msg_id!r} event_id={event_id!r} msg_seq={msg_seq}'
|
||||
)
|
||||
except Exception:
|
||||
if self.ap is not None:
|
||||
self.ap.logger.error(
|
||||
f'QQ Official: send_markdown_keyboard failed, falling back to text: {traceback.format_exc()}'
|
||||
)
|
||||
await self.logger.error(
|
||||
f'QQ Official: send_markdown_keyboard failed, falling back to text: {traceback.format_exc()}'
|
||||
)
|
||||
text_msg = platform_message.MessageChain([platform_message.Plain(text=plain_content)])
|
||||
try:
|
||||
await self.reply_message(message_source, text_msg)
|
||||
except Exception:
|
||||
pass
|
||||
return
|
||||
|
||||
sender_id = ''
|
||||
if source is not None:
|
||||
sender_id = (
|
||||
getattr(source, 'user_openid', None)
|
||||
or getattr(source, 'member_openid', None)
|
||||
or getattr(source, 'd_author_id', None)
|
||||
or ''
|
||||
)
|
||||
if not sender_id and message_source.sender is not None:
|
||||
sender_id = str(getattr(message_source.sender, 'id', '') or '')
|
||||
self._pending_forms[session_key] = {
|
||||
'form_data': form_data,
|
||||
'msg_id': msg_id,
|
||||
'sender_id': sender_id,
|
||||
'target_type': target_type,
|
||||
'target_id': target_id,
|
||||
'source_event_t': source.t if source is not None else None,
|
||||
'posted_at': time.time(),
|
||||
}
|
||||
await self.logger.info(
|
||||
f'QQ Official: form posted session={session_key} actions={len(form_data.get("actions") or [])}'
|
||||
)
|
||||
|
||||
async def _handle_interaction_create(
|
||||
self,
|
||||
event_data: dict,
|
||||
ws_event_id: typing.Optional[str] = None,
|
||||
) -> None:
|
||||
"""Handle a button-click INTERACTION_CREATE event.
|
||||
|
||||
Two IDs at play (QQ keeps them separate):
|
||||
ws_event_id top-level payload ``id`` (or webhook ``X-Bot-
|
||||
Event-Id``). The ONLY value accepted as
|
||||
``event_id`` for subsequent passive replies.
|
||||
d['id'] the interaction id — used for PUT
|
||||
/interactions/{id} ack. Cannot be reused as
|
||||
event_id (QQ returns 40034025 if you try).
|
||||
|
||||
Layout (https://bot.q.qq.com/.../msg-btn.html):
|
||||
chat_type 0 channel / 1 group / 2 c2c
|
||||
data.resolved.button_data what we set as ``action.data``
|
||||
data.resolved.button_id ``id`` field on the button row
|
||||
"""
|
||||
import langbot_plugin.api.entities.builtin.provider.session as provider_session
|
||||
|
||||
if self.ap is not None:
|
||||
self.ap.logger.info(
|
||||
f'QQ Official _handle_interaction_create entered; '
|
||||
f'ws_event_id={ws_event_id!r} '
|
||||
f'interaction_id={(event_data.get("id") if isinstance(event_data, dict) else None)!r} '
|
||||
f'chat_type={event_data.get("chat_type") if isinstance(event_data, dict) else None}'
|
||||
)
|
||||
|
||||
if not isinstance(event_data, dict):
|
||||
await self.logger.warning(f'QQ Official: INTERACTION_CREATE event_data is not dict: {type(event_data)}')
|
||||
return
|
||||
|
||||
# ACK uses the interaction id, NOT the ws event id.
|
||||
interaction_id = event_data.get('id') or ''
|
||||
if interaction_id:
|
||||
asyncio.create_task(self.bot.ack_interaction(interaction_id, code=0))
|
||||
|
||||
resolved = (event_data.get('data') or {}).get('resolved') or {}
|
||||
action_id = str(resolved.get('button_data') or resolved.get('button_id') or '').strip()
|
||||
if not action_id:
|
||||
await self.logger.warning('QQ Official: INTERACTION_CREATE missing button_data/button_id; ignoring')
|
||||
return
|
||||
|
||||
chat_type = event_data.get('chat_type')
|
||||
scene_target: typing.Optional[tuple[str, str]] = None
|
||||
if chat_type == 2 or event_data.get('user_openid'):
|
||||
scene_target = ('c2c', event_data.get('user_openid') or '')
|
||||
elif chat_type == 1 or event_data.get('group_openid'):
|
||||
scene_target = ('group', event_data.get('group_openid') or '')
|
||||
elif chat_type == 0 or event_data.get('channel_id'):
|
||||
scene_target = ('channel', event_data.get('channel_id') or '')
|
||||
|
||||
if not scene_target or not scene_target[1]:
|
||||
await self.logger.warning(f'QQ Official: INTERACTION_CREATE missing scene/target; raw={event_data}')
|
||||
return
|
||||
|
||||
target_type, target_id = scene_target
|
||||
session_key = f'{target_type}_{target_id}'
|
||||
|
||||
self._prune_pending_forms()
|
||||
pending = self._pending_forms.get(session_key)
|
||||
if not pending:
|
||||
await self.logger.warning(
|
||||
f'QQ Official: no pending form for session {session_key}; click ignored (action_id={action_id!r})'
|
||||
)
|
||||
return
|
||||
|
||||
# Cache ws_event_id so a follow-up pause / text reply can use it
|
||||
# as event_id for passive delivery (30-min window). Falls back to
|
||||
# the interaction_id only if no ws_event_id was provided (e.g.
|
||||
# tests / older payload shape) — QQ will reject that value but
|
||||
# we log so the mismatch is debuggable.
|
||||
cached_event_id = ws_event_id or interaction_id
|
||||
if cached_event_id:
|
||||
self._session_event_ids[session_key] = {
|
||||
'event_id': cached_event_id,
|
||||
'posted_at': time.time(),
|
||||
}
|
||||
# New anchor → fresh 5-reply budget.
|
||||
self._anchor_msg_seq[cached_event_id] = 0
|
||||
if self.ap is not None and not ws_event_id:
|
||||
self.ap.logger.warning(
|
||||
'QQ Official: INTERACTION_CREATE lacked ws_event_id; '
|
||||
'falling back to interaction_id (passive reply may be rejected)'
|
||||
)
|
||||
|
||||
form_data: dict = pending.get('form_data') or {}
|
||||
actions = form_data.get('actions') or []
|
||||
select_choice = resolve_select_button_action(form_data, action_id)
|
||||
if action_id.startswith(QQ_SELECT_ACTION_PREFIX) and select_choice is None:
|
||||
await self.logger.warning(f'QQ Official: invalid select action_id={action_id!r} for {session_key}')
|
||||
return
|
||||
|
||||
matched = None
|
||||
if select_choice is None:
|
||||
matched = next(
|
||||
(a for a in actions if str(a.get('id', '')) == action_id),
|
||||
None,
|
||||
)
|
||||
if matched is None:
|
||||
await self.logger.warning(
|
||||
f'QQ Official: action_id={action_id!r} is not present on pending form for {session_key}'
|
||||
)
|
||||
return
|
||||
self._pending_forms.pop(session_key, None)
|
||||
action_title = select_choice[1] if select_choice else matched.get('title') or action_id
|
||||
|
||||
initiator_id = str(pending.get('sender_id') or '')
|
||||
actor_id = str(event_data.get('member_openid') or event_data.get('user_openid') or initiator_id)
|
||||
|
||||
# Build resume payload matching the shape every other adapter uses
|
||||
# (DingTalk / Lark / Telegram / WeCom). The runner's
|
||||
# _merge_pending_form_action consumes this verbatim.
|
||||
if target_type == 'group' or target_type == 'channel':
|
||||
launcher_type = provider_session.LauncherTypes.GROUP
|
||||
launcher_id = target_id
|
||||
else:
|
||||
launcher_type = provider_session.LauncherTypes.PERSON
|
||||
launcher_id = target_id
|
||||
|
||||
form_action_data = {
|
||||
'form_token': form_data.get('form_token', ''),
|
||||
'workflow_run_id': form_data.get('workflow_run_id', ''),
|
||||
'action_id': '' if select_choice else action_id,
|
||||
'action_title': action_title,
|
||||
'node_title': form_data.get('node_title', ''),
|
||||
'user': f'{launcher_type.value}_{launcher_id}',
|
||||
'inputs': {'select': select_choice[1]} if select_choice else {},
|
||||
}
|
||||
if select_choice:
|
||||
form_action_data['_current_input_field'] = select_choice[0]
|
||||
form_action_data['_input_progress'] = True
|
||||
|
||||
event_label = 'Form Select' if select_choice else 'Form Action'
|
||||
message_chain = platform_message.MessageChain([platform_message.Plain(text=f'[{event_label}: {action_title}]')])
|
||||
|
||||
if launcher_type == provider_session.LauncherTypes.GROUP:
|
||||
synthetic_event: platform_events.MessageEvent = platform_events.GroupMessage(
|
||||
sender=platform_entities.GroupMember(
|
||||
id=actor_id or launcher_id,
|
||||
member_name='',
|
||||
permission='MEMBER',
|
||||
group=platform_entities.Group(
|
||||
id=launcher_id,
|
||||
name='',
|
||||
permission=platform_entities.Permission.Member,
|
||||
),
|
||||
special_title='',
|
||||
),
|
||||
message_chain=message_chain,
|
||||
time=int(time.time()),
|
||||
source_platform_object=None,
|
||||
)
|
||||
else:
|
||||
synthetic_event = platform_events.FriendMessage(
|
||||
sender=platform_entities.Friend(
|
||||
id=actor_id or launcher_id,
|
||||
nickname='',
|
||||
remark='',
|
||||
),
|
||||
message_chain=message_chain,
|
||||
time=int(time.time()),
|
||||
source_platform_object=None,
|
||||
)
|
||||
|
||||
if self.ap is None:
|
||||
await self.logger.error('QQ Official: ap not injected; cannot enqueue button-click query')
|
||||
return
|
||||
|
||||
bot_uuid = ''
|
||||
pipeline_uuid = form_data.get('pipeline_uuid') or None
|
||||
for bot in self.ap.platform_mgr.bots:
|
||||
if bot.adapter is self:
|
||||
bot_uuid = bot.bot_entity.uuid
|
||||
pipeline_uuid = pipeline_uuid or bot.bot_entity.use_pipeline_uuid
|
||||
break
|
||||
|
||||
try:
|
||||
await self.ap.query_pool.add_query(
|
||||
bot_uuid=bot_uuid,
|
||||
launcher_type=launcher_type,
|
||||
launcher_id=launcher_id,
|
||||
sender_id=actor_id or launcher_id,
|
||||
message_event=synthetic_event,
|
||||
message_chain=message_chain,
|
||||
adapter=self,
|
||||
pipeline_uuid=pipeline_uuid,
|
||||
variables={
|
||||
'_dify_form_action': form_action_data,
|
||||
'_routed_by_rule': True,
|
||||
},
|
||||
)
|
||||
await self.logger.info(
|
||||
f'QQ Official: button-click query enqueued action_id={action_id!r} '
|
||||
f'session={session_key} actor_id={actor_id}'
|
||||
)
|
||||
except Exception:
|
||||
await self.logger.error(f'QQ Official: enqueue button-click query failed: {traceback.format_exc()}')
|
||||
|
||||
@@ -31,18 +31,6 @@ spec:
|
||||
type: array[string]
|
||||
required: false
|
||||
default: []
|
||||
- name: one-click-bind
|
||||
label:
|
||||
en_US: One-Click QR Binding
|
||||
zh_Hans: 一键扫码绑定
|
||||
zh_Hant: 一鍵掃碼綁定
|
||||
description:
|
||||
en_US: Scan QR code with mobile QQ to auto-fill AppID and Secret (Token is not used and can be left blank)
|
||||
zh_Hans: 使用手机 QQ 扫码绑定,自动填写 AppID 和密钥(当前未使用 Token,可留空)
|
||||
zh_Hant: 使用手機 QQ 掃碼綁定,自動填寫 AppID 和密鑰(目前未使用 Token,可留空)
|
||||
type: qr-code-login
|
||||
login_platform: qqofficial
|
||||
required: false
|
||||
- name: appid
|
||||
label:
|
||||
en_US: App ID
|
||||
@@ -64,12 +52,8 @@ spec:
|
||||
en_US: Token
|
||||
zh_Hans: 令牌
|
||||
zh_Hant: 令牌
|
||||
description:
|
||||
en_US: Optional. The QR binding cannot return this value; the current adapter implementation does not use it either, so it can be safely left blank.
|
||||
zh_Hans: 可选。扫码绑定无法获取该字段,当前适配器实现也未使用该字段,留空即可。
|
||||
zh_Hant: 可選。掃碼綁定無法取得此欄位,目前介面卡實作亦未使用,留空即可。
|
||||
type: string
|
||||
required: false
|
||||
required: true
|
||||
default: ""
|
||||
- name: enable-webhook
|
||||
label:
|
||||
|
||||
@@ -1,17 +1,15 @@
|
||||
from __future__ import annotations
|
||||
import time
|
||||
|
||||
|
||||
import telegram
|
||||
import telegram.ext
|
||||
from telegram import ForceReply, InlineKeyboardButton, InlineKeyboardMarkup, Update
|
||||
from telegram.ext import ApplicationBuilder, ContextTypes, MessageHandler, CallbackQueryHandler, filters
|
||||
from telegram import Update
|
||||
from telegram.ext import ApplicationBuilder, ContextTypes, MessageHandler, filters
|
||||
import telegramify_markdown
|
||||
import typing
|
||||
import traceback
|
||||
import json
|
||||
import base64
|
||||
import time
|
||||
import uuid
|
||||
import pydantic
|
||||
|
||||
from langbot.pkg.utils import httpclient
|
||||
@@ -22,61 +20,6 @@ import langbot_plugin.api.entities.builtin.platform.entities as platform_entitie
|
||||
import langbot_plugin.api.definition.abstract.platform.event_logger as abstract_platform_logger
|
||||
|
||||
|
||||
def _telegram_select_field_options(form_data: dict) -> tuple[str, list[str]]:
|
||||
"""Return the active select field and its option values."""
|
||||
field_name = str(form_data.get('_current_input_field') or '').strip()
|
||||
if not field_name:
|
||||
return '', []
|
||||
field = next(
|
||||
(
|
||||
item
|
||||
for item in form_data.get('input_defs') or []
|
||||
if str(item.get('output_variable_name') or '').strip() == field_name
|
||||
),
|
||||
None,
|
||||
)
|
||||
if not field or str(field.get('type') or '').strip().lower() != 'select':
|
||||
return '', []
|
||||
|
||||
source = field.get('option_source') or {}
|
||||
source_value = source.get('value') if isinstance(source, dict) else None
|
||||
if isinstance(source_value, list):
|
||||
return field_name, [str(item) for item in source_value]
|
||||
if isinstance(source_value, str):
|
||||
return field_name, [part.strip() for part in source_value.splitlines() if part.strip()]
|
||||
|
||||
options = field.get('options')
|
||||
if not isinstance(options, list):
|
||||
return field_name, []
|
||||
values = []
|
||||
for item in options:
|
||||
if isinstance(item, dict):
|
||||
values.append(str(item.get('label') or item.get('value') or ''))
|
||||
else:
|
||||
values.append(str(item))
|
||||
return field_name, [value for value in values if value]
|
||||
|
||||
|
||||
def _telegram_form_action_from_callback(data: dict) -> dict | None:
|
||||
"""Translate compact Telegram callback data into a runner form action."""
|
||||
if 'x' not in data:
|
||||
return {
|
||||
'action_id': str(data.get('action_id') or data.get('a') or ''),
|
||||
'inputs': {},
|
||||
}
|
||||
try:
|
||||
option_index = int(data['x'])
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if option_index < 0:
|
||||
return None
|
||||
return {
|
||||
'action_id': '',
|
||||
'inputs': {'select': {'index': option_index}},
|
||||
'_input_progress': True,
|
||||
}
|
||||
|
||||
|
||||
class TelegramMessageConverter(abstract_platform_adapter.AbstractMessageConverter):
|
||||
@staticmethod
|
||||
async def yiri2target(message_chain: platform_message.MessageChain, bot: telegram.Bot) -> list[dict]:
|
||||
@@ -224,7 +167,7 @@ class TelegramEventConverter(abstract_platform_adapter.AbstractEventConverter):
|
||||
time=event.message.date.timestamp(),
|
||||
source_platform_object=event,
|
||||
)
|
||||
elif event.effective_chat.type in ('group', 'supergroup'):
|
||||
elif event.effective_chat.type == 'group' or 'supergroup':
|
||||
return platform_events.GroupMessage(
|
||||
sender=platform_entities.GroupMember(
|
||||
id=event.effective_chat.id,
|
||||
@@ -246,7 +189,6 @@ class TelegramEventConverter(abstract_platform_adapter.AbstractEventConverter):
|
||||
class TelegramAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
bot: telegram.Bot = pydantic.Field(exclude=True)
|
||||
application: telegram.ext.Application = pydantic.Field(exclude=True)
|
||||
ap: typing.Any = pydantic.Field(exclude=True, default=None)
|
||||
|
||||
message_converter: TelegramMessageConverter = TelegramMessageConverter()
|
||||
event_converter: TelegramEventConverter = TelegramEventConverter()
|
||||
@@ -262,48 +204,6 @@ class TelegramAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
typing.Callable[[platform_events.Event, abstract_platform_adapter.AbstractMessagePlatformAdapter], None],
|
||||
] = {}
|
||||
|
||||
_FORM_ACTION_CACHE_TTL = 30 * 60
|
||||
# callback_data -> (display title, pipeline UUID, expiration time, form group id)
|
||||
_form_action_titles: typing.Dict[str, tuple[str, str, float, str]] = {}
|
||||
|
||||
def _prune_form_action_titles(self, now: float | None = None) -> None:
|
||||
now = time.monotonic() if now is None else now
|
||||
expired = [key for key, (_, _, expires_at, _) in self._form_action_titles.items() if expires_at <= now]
|
||||
for key in expired:
|
||||
self._form_action_titles.pop(key, None)
|
||||
|
||||
def _cache_form_action_titles(
|
||||
self,
|
||||
mappings: dict[str, str],
|
||||
pipeline_uuid: str = '',
|
||||
now: float | None = None,
|
||||
) -> None:
|
||||
now = time.monotonic() if now is None else now
|
||||
self._prune_form_action_titles(now)
|
||||
group_id = uuid.uuid4().hex
|
||||
expires_at = now + self._FORM_ACTION_CACHE_TTL
|
||||
self._form_action_titles.update(
|
||||
{callback_data: (title, pipeline_uuid, expires_at, group_id) for callback_data, title in mappings.items()}
|
||||
)
|
||||
|
||||
def _take_form_action_context(self, callback_data: str, now: float | None = None) -> tuple[str, str] | None:
|
||||
"""Consume a callback and invalidate every button from the same form."""
|
||||
self._prune_form_action_titles(now)
|
||||
entry = self._form_action_titles.get(callback_data)
|
||||
if entry is None:
|
||||
return None
|
||||
title, pipeline_uuid, _, group_id = entry
|
||||
group_keys = [
|
||||
key for key, (_, _, _, cached_group_id) in self._form_action_titles.items() if cached_group_id == group_id
|
||||
]
|
||||
for key in group_keys:
|
||||
self._form_action_titles.pop(key, None)
|
||||
return title, pipeline_uuid
|
||||
|
||||
def _take_form_action_title(self, callback_data: str, now: float | None = None) -> str | None:
|
||||
context = self._take_form_action_context(callback_data, now)
|
||||
return context[0] if context else None
|
||||
|
||||
def __init__(self, config: dict, logger: abstract_platform_logger.AbstractEventLogger):
|
||||
async def telegram_callback(update: Update, context: ContextTypes.DEFAULT_TYPE):
|
||||
if update.message.from_user.is_bot:
|
||||
@@ -324,117 +224,6 @@ class TelegramAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
telegram_callback,
|
||||
)
|
||||
)
|
||||
|
||||
async def callback_query_handler(update: Update, context: ContextTypes.DEFAULT_TYPE):
|
||||
query = update.callback_query
|
||||
await query.answer()
|
||||
try:
|
||||
data = json.loads(query.data)
|
||||
if data.get('form_action') or data.get('f'):
|
||||
import langbot_plugin.api.entities.builtin.provider.session as provider_session
|
||||
|
||||
# workflow_run_id is not in the callback payload (too large
|
||||
# for Telegram's 64-byte limit). Only w_suffix is sent;
|
||||
# the runner resolves the full run id from _PENDING_FORMS.
|
||||
w_suffix = data.get('w', '')
|
||||
session_key = data.get('session_key') or data.get('s', '')
|
||||
callback_action = _telegram_form_action_from_callback(data)
|
||||
action_context = self._take_form_action_context(query.data) if callback_action is not None else None
|
||||
if callback_action is None or action_context is None:
|
||||
await self.logger.warning(f'Invalid or stale Telegram form callback: {query.data!r}')
|
||||
return
|
||||
action_title, pipeline_uuid = action_context
|
||||
# Show selected action feedback by editing the original message
|
||||
try:
|
||||
original_text = query.message.text or ''
|
||||
selected_text = f'{original_text}\n\n✅ {action_title}'
|
||||
await query.edit_message_text(text=selected_text, reply_markup=None)
|
||||
except Exception:
|
||||
# If edit fails (e.g. message too long), just pass
|
||||
pass
|
||||
|
||||
if session_key.startswith('group_') or session_key.startswith('g:'):
|
||||
launcher_type = provider_session.LauncherTypes.GROUP
|
||||
launcher_id = (
|
||||
session_key.split(':', 1)[1]
|
||||
if session_key.startswith('g:')
|
||||
else session_key[len('group_') :]
|
||||
)
|
||||
else:
|
||||
launcher_type = provider_session.LauncherTypes.PERSON
|
||||
launcher_id = (
|
||||
session_key.split(':', 1)[1]
|
||||
if session_key.startswith('p:')
|
||||
else session_key[len('person_') :]
|
||||
)
|
||||
|
||||
user_id = str(query.from_user.id)
|
||||
|
||||
# Find bot_uuid and pipeline_uuid
|
||||
bot_uuid = ''
|
||||
for b in self.ap.platform_mgr.bots:
|
||||
if b.adapter is self:
|
||||
bot_uuid = b.bot_entity.uuid
|
||||
pipeline_uuid = pipeline_uuid or b.bot_entity.use_pipeline_uuid
|
||||
break
|
||||
|
||||
form_action_data = {
|
||||
# workflow_run_id is intentionally omitted; the runner
|
||||
# resolves it from w_suffix via _PENDING_FORMS.
|
||||
'w_suffix': w_suffix,
|
||||
'user': f'{launcher_type.value}_{launcher_id}',
|
||||
**callback_action,
|
||||
}
|
||||
|
||||
event_label = 'Form Select' if callback_action.get('_input_progress') else 'Form Action'
|
||||
message_chain = platform_message.MessageChain(
|
||||
[platform_message.Plain(text=f'[{event_label}: {action_title}]')]
|
||||
)
|
||||
|
||||
if launcher_type == provider_session.LauncherTypes.GROUP:
|
||||
synthetic_event = platform_events.GroupMessage(
|
||||
sender=platform_entities.GroupMember(
|
||||
id=user_id,
|
||||
member_name='',
|
||||
permission=platform_entities.Permission.Member,
|
||||
group=platform_entities.Group(
|
||||
id=launcher_id,
|
||||
name='',
|
||||
permission=platform_entities.Permission.Member,
|
||||
),
|
||||
),
|
||||
message_chain=message_chain,
|
||||
source_platform_object=update,
|
||||
)
|
||||
else:
|
||||
synthetic_event = platform_events.FriendMessage(
|
||||
sender=platform_entities.Friend(
|
||||
id=user_id,
|
||||
nickname='',
|
||||
remark='',
|
||||
),
|
||||
message_chain=message_chain,
|
||||
source_platform_object=update,
|
||||
)
|
||||
|
||||
await self.ap.query_pool.add_query(
|
||||
bot_uuid=bot_uuid,
|
||||
launcher_type=launcher_type,
|
||||
launcher_id=launcher_id,
|
||||
sender_id=user_id,
|
||||
message_event=synthetic_event,
|
||||
message_chain=message_chain,
|
||||
adapter=self,
|
||||
pipeline_uuid=pipeline_uuid,
|
||||
variables={
|
||||
'_dify_form_action': form_action_data,
|
||||
'_routed_by_rule': True,
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
await self.logger.error(f'Error in telegram callback query: {traceback.format_exc()}')
|
||||
|
||||
application.add_handler(CallbackQueryHandler(callback_query_handler))
|
||||
super().__init__(
|
||||
config=config,
|
||||
logger=logger,
|
||||
@@ -525,34 +314,23 @@ class TelegramAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
args['parse_mode'] = 'MarkdownV2'
|
||||
return args
|
||||
|
||||
async def _delete_group_stream_message(self, chat_mode: str, chat_id: int, stream_id: int | None):
|
||||
if chat_mode != 'group' or stream_id is None:
|
||||
return
|
||||
try:
|
||||
await self.bot.delete_message(chat_id=chat_id, message_id=stream_id)
|
||||
except telegram.error.TelegramError:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _is_form_placeholder_chunk(text: str) -> bool:
|
||||
"""Return True for invisible placeholder chunks used to carry forms."""
|
||||
|
||||
if not text:
|
||||
return True
|
||||
|
||||
cleaned = text.replace('\u200b', '').replace('\u200c', '').replace('\u200d', '').replace('\ufeff', '').strip()
|
||||
return cleaned == ''
|
||||
|
||||
async def create_message_card(self, message_id, event):
|
||||
assert isinstance(event.source_platform_object, Update)
|
||||
update = event.source_platform_object
|
||||
chat_id = update.effective_chat.id
|
||||
effective_message = update.effective_message
|
||||
message_thread_id = getattr(effective_message, 'message_thread_id', None) if effective_message else None
|
||||
chat_type = update.effective_chat.type
|
||||
message_thread_id = update.message.message_thread_id
|
||||
|
||||
args = self._build_message_args(chat_id, 'Thinking...', message_thread_id)
|
||||
send_msg = await self.bot.send_message(**args)
|
||||
self.msg_stream_id[message_id] = ('message', send_msg.message_id, False)
|
||||
if chat_type == 'private':
|
||||
draft_id = int(time.time() * 1000)
|
||||
self.msg_stream_id[message_id] = ('private', draft_id)
|
||||
|
||||
args = self._build_message_args(chat_id, 'Thinking...', message_thread_id, draft_id=draft_id)
|
||||
await self.bot.send_message_draft(**args)
|
||||
else:
|
||||
args = self._build_message_args(chat_id, 'Thinking...', message_thread_id)
|
||||
send_msg = await self.bot.send_message(**args)
|
||||
self.msg_stream_id[message_id] = ('group', send_msg.message_id)
|
||||
|
||||
return True
|
||||
|
||||
@@ -569,15 +347,12 @@ class TelegramAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
assert isinstance(message_source.source_platform_object, Update)
|
||||
update = message_source.source_platform_object
|
||||
chat_id = update.effective_chat.id
|
||||
effective_message = update.effective_message
|
||||
message_thread_id = getattr(effective_message, 'message_thread_id', None) if effective_message else None
|
||||
message_thread_id = update.message.message_thread_id
|
||||
|
||||
if message_id not in self.msg_stream_id:
|
||||
return
|
||||
|
||||
stream_state = self.msg_stream_id[message_id]
|
||||
chat_mode, stream_id = stream_state[:2]
|
||||
has_visible_content = len(stream_state) > 2 and stream_state[2]
|
||||
chat_mode, draft_id = self.msg_stream_id[message_id]
|
||||
components = await TelegramMessageConverter.yiri2target(message, self.bot)
|
||||
|
||||
if not components or components[0]['type'] != 'text':
|
||||
@@ -586,68 +361,17 @@ class TelegramAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
return
|
||||
|
||||
content = components[0]['text']
|
||||
form_data = getattr(bot_message, '_form_data', None)
|
||||
|
||||
if form_data and is_final:
|
||||
if not has_visible_content:
|
||||
await self._send_form_action_buttons(message_source, form_data, edit_message_id=stream_id)
|
||||
else:
|
||||
await self._send_form_action_buttons(message_source, form_data)
|
||||
self.msg_stream_id.pop(message_id, None)
|
||||
return
|
||||
|
||||
if self._is_form_placeholder_chunk(content):
|
||||
if is_final and bot_message.tool_calls is None and not has_visible_content:
|
||||
await self._delete_group_stream_message(chat_mode, chat_id, stream_id)
|
||||
self.msg_stream_id.pop(message_id, None)
|
||||
return
|
||||
|
||||
if chat_mode == 'private':
|
||||
# Streaming via draft (ephemeral preview in the chat input area)
|
||||
if (msg_seq - 1) % 8 == 0 or is_final:
|
||||
args = self._build_message_args(chat_id, content, message_thread_id, draft_id=stream_id)
|
||||
try:
|
||||
await self.bot.send_message_draft(**args)
|
||||
except telegram.error.BadRequest as exc:
|
||||
if 'Message_too_long' in str(exc):
|
||||
args['text'] = content[:4000] + '\n\n… (truncated)'
|
||||
try:
|
||||
await self.bot.send_message_draft(**args)
|
||||
except telegram.error.RetryAfter:
|
||||
pass
|
||||
else:
|
||||
pass # Ignore other draft errors (cosmetic)
|
||||
self.msg_stream_id[message_id] = (chat_mode, stream_id, True)
|
||||
args = self._build_message_args(chat_id, content, message_thread_id, draft_id=draft_id)
|
||||
await self.bot.send_message_draft(**args)
|
||||
if is_final and bot_message.tool_calls is None:
|
||||
# Finalise: send the real message, discard the draft
|
||||
args = self._build_message_args(chat_id, content, message_thread_id)
|
||||
try:
|
||||
await self.bot.send_message(**args)
|
||||
except telegram.error.BadRequest as exc:
|
||||
if 'Message_too_long' in str(exc):
|
||||
args['text'] = content[:4000] + '\n\n… (truncated)'
|
||||
await self.bot.send_message(**args)
|
||||
else:
|
||||
raise
|
||||
del args['draft_id']
|
||||
await self.bot.send_message(**args)
|
||||
self.msg_stream_id.pop(message_id)
|
||||
else:
|
||||
# Streaming via edit_message_text (persistent message)
|
||||
if stream_id is None:
|
||||
args = self._build_message_args(chat_id, content, message_thread_id)
|
||||
try:
|
||||
send_msg = await self.bot.send_message(**args)
|
||||
except telegram.error.BadRequest as exc:
|
||||
if 'Message_too_long' in str(exc):
|
||||
args['text'] = self._process_markdown(content[:4000] + '\n\n鈥?(truncated)')
|
||||
send_msg = await self.bot.send_message(**args)
|
||||
else:
|
||||
raise
|
||||
self.msg_stream_id[message_id] = (chat_mode, send_msg.message_id, True)
|
||||
if is_final and bot_message.tool_calls is None:
|
||||
self.msg_stream_id.pop(message_id, None)
|
||||
return
|
||||
|
||||
if not has_visible_content or (msg_seq - 1) % 8 == 0 or is_final:
|
||||
stream_id = draft_id
|
||||
if (msg_seq - 1) % 8 == 0 or is_final:
|
||||
args = {
|
||||
'message_id': stream_id,
|
||||
'chat_id': chat_id,
|
||||
@@ -655,137 +379,11 @@ class TelegramAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
}
|
||||
if self.config.get('markdown_card', False):
|
||||
args['parse_mode'] = 'MarkdownV2'
|
||||
try:
|
||||
await self.bot.edit_message_text(**args)
|
||||
except telegram.error.BadRequest as exc:
|
||||
if 'Message_too_long' in str(exc):
|
||||
args['text'] = self._process_markdown(content[:4000] + '\n\n… (truncated)')
|
||||
await self.bot.edit_message_text(**args)
|
||||
else:
|
||||
raise
|
||||
self.msg_stream_id[message_id] = (chat_mode, stream_id, True)
|
||||
await self.bot.edit_message_text(**args)
|
||||
|
||||
if is_final and bot_message.tool_calls is None:
|
||||
self.msg_stream_id.pop(message_id)
|
||||
|
||||
async def _send_form_action_buttons(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
form_data: dict,
|
||||
edit_message_id: int | None = None,
|
||||
):
|
||||
"""Send inline keyboard buttons for Dify form fields or actions."""
|
||||
actions = form_data.get('actions', [])
|
||||
node_title = form_data.get('node_title', '')
|
||||
form_content = form_data.get('form_content', '')
|
||||
workflow_run_id = form_data.get('workflow_run_id', '')
|
||||
# Telegram callback_data is capped at 64 bytes, so we identify the
|
||||
# paused workflow by the last 8 chars of workflow_run_id (unique
|
||||
# within a session with overwhelming probability).
|
||||
w_suffix = workflow_run_id[-8:] if workflow_run_id else ''
|
||||
|
||||
if isinstance(message_source, platform_events.GroupMessage):
|
||||
session_key = f'g:{message_source.group.id}'
|
||||
else:
|
||||
session_key = f'p:{message_source.sender.id}'
|
||||
|
||||
current_field = str(form_data.get('_current_input_field') or '').strip()
|
||||
is_field_step = bool(current_field) and not form_data.get('_action_select_only')
|
||||
select_field, select_options = _telegram_select_field_options(form_data)
|
||||
is_select_field = bool(select_field and select_options)
|
||||
if is_select_field:
|
||||
choices = [(option, {'x': idx}) for idx, option in enumerate(select_options)]
|
||||
elif is_field_step:
|
||||
choices = []
|
||||
else:
|
||||
choices = [(action.get('title', action.get('id', '')), {'a': action.get('id', '')}) for action in actions]
|
||||
|
||||
keyboard = []
|
||||
pending_title_mappings: dict[str, str] = {}
|
||||
oversized = False
|
||||
buttons_per_row = 2 if is_select_field else 1
|
||||
current_row = []
|
||||
for title, choice_data in choices:
|
||||
callback_payload = {'f': 1, **choice_data, 's': session_key}
|
||||
if w_suffix:
|
||||
callback_payload['w'] = w_suffix
|
||||
callback_data = json.dumps(callback_payload, separators=(',', ':'))
|
||||
if len(callback_data.encode('utf-8')) > 64:
|
||||
oversized = True
|
||||
break
|
||||
pending_title_mappings[callback_data] = str(title)
|
||||
current_row.append(InlineKeyboardButton(str(title), callback_data=callback_data))
|
||||
if len(current_row) == buttons_per_row:
|
||||
keyboard.append(current_row)
|
||||
current_row = []
|
||||
if current_row and not oversized:
|
||||
keyboard.append(current_row)
|
||||
|
||||
update = message_source.source_platform_object
|
||||
chat_id = update.effective_chat.id
|
||||
effective_message = update.effective_message
|
||||
message_thread_id = getattr(effective_message, 'message_thread_id', None) if effective_message else None
|
||||
|
||||
heading = f'[{node_title}]'
|
||||
text_lines = [heading]
|
||||
if form_content:
|
||||
text_lines.append(form_content)
|
||||
|
||||
if oversized:
|
||||
# callback_data exceeds Telegram's 64-byte limit — fall back to
|
||||
# a plain-text numbered list so the user can reply by number.
|
||||
for idx, (title, _) in enumerate(choices, start=1):
|
||||
text_lines.append(f' {idx}. {title}')
|
||||
args = {
|
||||
'chat_id': chat_id,
|
||||
'text': '\n\n'.join(text_lines),
|
||||
}
|
||||
elif keyboard:
|
||||
self._cache_form_action_titles(
|
||||
pending_title_mappings,
|
||||
str(form_data.get('pipeline_uuid') or ''),
|
||||
)
|
||||
reply_markup = InlineKeyboardMarkup(keyboard)
|
||||
args = {
|
||||
'chat_id': chat_id,
|
||||
'text': '\n\n'.join(text_lines),
|
||||
'reply_markup': reply_markup,
|
||||
}
|
||||
elif is_field_step:
|
||||
args = {
|
||||
'chat_id': chat_id,
|
||||
'text': '\n\n'.join(text_lines),
|
||||
# Telegram privacy-mode bots receive replies to ForceReply
|
||||
# prompts even when they cannot read ordinary group messages.
|
||||
'reply_markup': ForceReply(
|
||||
selective=False,
|
||||
input_field_placeholder=current_field,
|
||||
),
|
||||
}
|
||||
else:
|
||||
args = {
|
||||
'chat_id': chat_id,
|
||||
'text': '\n\n'.join(text_lines),
|
||||
}
|
||||
|
||||
if message_thread_id:
|
||||
args['message_thread_id'] = message_thread_id
|
||||
|
||||
if edit_message_id is not None:
|
||||
edit_args = {
|
||||
'chat_id': chat_id,
|
||||
'message_id': edit_message_id,
|
||||
'text': args['text'],
|
||||
}
|
||||
edit_args['reply_markup'] = args.get('reply_markup')
|
||||
try:
|
||||
await self.bot.edit_message_text(**edit_args)
|
||||
return
|
||||
except telegram.error.TelegramError:
|
||||
await self._delete_group_stream_message('group', chat_id, edit_message_id)
|
||||
|
||||
await self.bot.send_message(**args)
|
||||
|
||||
def get_launcher_id(self, event: platform_events.MessageEvent) -> str | None:
|
||||
if not isinstance(event.source_platform_object, Update):
|
||||
return None
|
||||
|
||||
@@ -422,6 +422,64 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
session_type=session_type,
|
||||
)
|
||||
|
||||
# Determine if pipeline_uuid is a workflow or a legacy pipeline by querying both services
|
||||
workflow_dict = await self.ap.workflow_service.get_workflow(pipeline_uuid)
|
||||
pipeline_dict = await self.ap.pipeline_service.get_pipeline(pipeline_uuid)
|
||||
|
||||
if workflow_dict is not None:
|
||||
# UUID exists in workflow table - execute as workflow
|
||||
# Set pipeline_uuid for workflow nodes to broadcast messages correctly
|
||||
self.ap.platform_mgr.websocket_proxy_bot.bot_entity.use_pipeline_uuid = pipeline_uuid
|
||||
|
||||
message_content = str(message_chain)
|
||||
message_context = {
|
||||
'message_id': str(message_id),
|
||||
'message_content': message_content,
|
||||
'sender_id': f'websocket_{connection.connection_id}',
|
||||
'sender_name': 'User',
|
||||
'platform': 'websocket',
|
||||
'conversation_id': connection.connection_id,
|
||||
'is_group': session_type == 'group',
|
||||
'group_id': 'websocketgroup' if session_type == 'group' else None,
|
||||
'mentions': [],
|
||||
'reply_to': None,
|
||||
'raw_message': {
|
||||
'message': message_chain_obj,
|
||||
'connection_id': connection.connection_id,
|
||||
'session_type': session_type,
|
||||
},
|
||||
}
|
||||
|
||||
trigger_data = {
|
||||
'message': message_content,
|
||||
'message_chain': message_chain_obj,
|
||||
'session_type': session_type,
|
||||
'connection_id': connection.connection_id,
|
||||
'message_context': message_context,
|
||||
}
|
||||
|
||||
try:
|
||||
from ...api.http.service.workflow import WorkflowExecutionFailedError
|
||||
|
||||
# Log workflow execution start (matching pipeline logging)
|
||||
session_id = f'{session_type}_{connection.connection_id}'
|
||||
logger.info(f'Processing workflow message from {session_id}: {message_content}')
|
||||
|
||||
execution_id = await self.ap.workflow_service.execute_workflow(
|
||||
pipeline_uuid, # This is actually a workflow UUID
|
||||
trigger_type='message',
|
||||
trigger_data=trigger_data,
|
||||
session_id=session_id,
|
||||
user_id=message_context['sender_id'],
|
||||
bot_id=self.ap.platform_mgr.websocket_proxy_bot.bot_entity.uuid,
|
||||
)
|
||||
except WorkflowExecutionFailedError as e:
|
||||
await connection.send_queue.put({'type': 'error', 'message': e.message})
|
||||
except Exception as e:
|
||||
logger.error(f'Workflow websocket execution error: {e}', exc_info=True)
|
||||
await connection.send_queue.put({'type': 'error', 'message': str(e)})
|
||||
return
|
||||
|
||||
# 添加消息源
|
||||
message_chain.insert(0, platform_message.Source(id=message_id, time=datetime.now().timestamp()))
|
||||
|
||||
|
||||
@@ -11,13 +11,7 @@ import langbot_plugin.api.entities.builtin.platform.events as platform_events
|
||||
import langbot_plugin.api.entities.builtin.platform.entities as platform_entities
|
||||
from ..logger import EventLogger
|
||||
from langbot.libs.wecom_ai_bot_api.wecombotevent import WecomBotEvent
|
||||
from langbot.libs.wecom_ai_bot_api.api import (
|
||||
WecomBotClient,
|
||||
extract_template_card_action,
|
||||
extract_template_card_event_payload,
|
||||
extract_template_card_selections,
|
||||
parse_select_button_action,
|
||||
)
|
||||
from langbot.libs.wecom_ai_bot_api.api import WecomBotClient
|
||||
from langbot.libs.wecom_ai_bot_api.ws_client import WecomBotWsClient
|
||||
|
||||
|
||||
@@ -302,7 +296,6 @@ class WecomBotAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
listeners: dict = {}
|
||||
_stream_to_monitoring_msg: dict = {} # Maps stream_id to (monitoring_message_id, timestamp)
|
||||
_STREAM_MAPPING_TTL = 600 # 10 minutes
|
||||
ap: typing.Any = None
|
||||
|
||||
def __init__(self, config: dict, logger: EventLogger):
|
||||
enable_webhook = config.get('enable-webhook', False)
|
||||
@@ -343,25 +336,6 @@ class WecomBotAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
_stream_to_monitoring_msg={},
|
||||
)
|
||||
|
||||
# Both WecomBotClient (webhook) and WecomBotWsClient (ws long-conn)
|
||||
# expose ``set_card_action_callback``. Wire the click handler so
|
||||
# Dify human-input button taps resume the workflow on either mode.
|
||||
if hasattr(self.bot, 'set_card_action_callback'):
|
||||
self.bot.set_card_action_callback(self._on_card_action)
|
||||
|
||||
# Hand the client a `source` block so every interactive
|
||||
# template_card it emits carries the LangBot logo + name at the
|
||||
# top — the WeCom analogue of DingTalk's Avatar header.
|
||||
# Always on; icon_url accepts plain HTTPS URLs (no upload needed).
|
||||
if hasattr(self.bot, 'set_card_source'):
|
||||
self.bot.set_card_source(
|
||||
{
|
||||
'icon_url': 'https://raw.githubusercontent.com/RockChinQ/LangBot/master/res/logo-blue.png',
|
||||
'desc': 'LangBot',
|
||||
'desc_color': 0,
|
||||
}
|
||||
)
|
||||
|
||||
async def reply_message(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
@@ -371,37 +345,15 @@ class WecomBotAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
content = await self.message_converter.yiri2target(message)
|
||||
_ws_mode = not self.config.get('enable-webhook', False)
|
||||
|
||||
event = message_source.source_platform_object
|
||||
# Synthetic events (button-click resume queries) have no inbound
|
||||
# platform object. Fall back to a proactive send so error
|
||||
# messages and one-shot replies still reach the user.
|
||||
if event is None:
|
||||
if _ws_mode:
|
||||
if isinstance(message_source, platform_events.GroupMessage):
|
||||
chat_id = str(message_source.group.id)
|
||||
else:
|
||||
chat_id = str(message_source.sender.id)
|
||||
try:
|
||||
await self.bot.send_message(chat_id, content)
|
||||
except Exception:
|
||||
await self.logger.error(
|
||||
f'WeComBot: proactive reply for synthetic event failed: {traceback.format_exc()}'
|
||||
)
|
||||
else:
|
||||
await self.logger.warning(
|
||||
'WeComBot webhook mode cannot reply to a synthetic event '
|
||||
'(no req_id and no proactive-send credentials); dropping.'
|
||||
)
|
||||
return
|
||||
|
||||
if _ws_mode:
|
||||
req_id = event.get('req_id', '') if isinstance(event, dict) else getattr(event, 'req_id', '')
|
||||
event = message_source.source_platform_object
|
||||
req_id = event.get('req_id', '')
|
||||
if req_id:
|
||||
await self.bot.reply_text(req_id, content)
|
||||
else:
|
||||
await self.bot.set_message(event.message_id, content)
|
||||
else:
|
||||
await self.bot.set_message(event.message_id, content)
|
||||
await self.bot.set_message(message_source.source_platform_object.message_id, content)
|
||||
|
||||
async def reply_message_chunk(
|
||||
self,
|
||||
@@ -412,56 +364,9 @@ class WecomBotAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
is_final: bool = False,
|
||||
):
|
||||
content = await self.message_converter.yiri2target(message)
|
||||
msg_id = message_source.source_platform_object.message_id
|
||||
_ws_mode = not self.config.get('enable-webhook', False)
|
||||
|
||||
# Synthetic events (e.g. button-click triggered form resume) have
|
||||
# no inbound platform message — no msg_id, no req_id, no stream
|
||||
# session. The output must go via the proactive-send path instead
|
||||
# of the stream/reply path.
|
||||
spo = message_source.source_platform_object
|
||||
if spo is None:
|
||||
return await self._handle_synthetic_chunk(message_source, bot_message, content, is_final, _ws_mode)
|
||||
|
||||
msg_id = spo.message_id
|
||||
|
||||
# Dify human-input pause: when the runner attaches `_form_data` to
|
||||
# the final chunk, hand the button_interaction card off to the
|
||||
# underlying client. In webhook mode the card is queued for the
|
||||
# next followup poll; in ws mode it's sent as a reply frame
|
||||
# immediately. Falls back to plain text when the bot has no active
|
||||
# stream session for this msg_id (rare).
|
||||
form_data = getattr(bot_message, '_form_data', None)
|
||||
if form_data and is_final:
|
||||
if hasattr(self.bot, 'push_form_pause'):
|
||||
ok, stream_id, task_id = await self.bot.push_form_pause(msg_id, form_data)
|
||||
if ok:
|
||||
await self.logger.info(
|
||||
f'WeComBot: pending button_interaction registered '
|
||||
f'stream_id={stream_id} task_id={task_id} ws_mode={_ws_mode}'
|
||||
)
|
||||
return {'stream': True, 'form': True, 'task_id': task_id}
|
||||
await self.logger.warning(
|
||||
'WeComBot: cannot register form pause (no active stream session); falling back to plain text'
|
||||
)
|
||||
try:
|
||||
from langbot.pkg.provider.runners.difysvapi import _format_human_input_text
|
||||
|
||||
fallback = _format_human_input_text(
|
||||
form_data.get('node_title', ''),
|
||||
form_data.get('form_content', ''),
|
||||
form_data.get('actions', []) or [],
|
||||
)
|
||||
except Exception:
|
||||
fallback = content or '(人工输入)'
|
||||
if _ws_mode:
|
||||
event = message_source.source_platform_object
|
||||
req_id = event.get('req_id', '') if isinstance(event, dict) else getattr(event, 'req_id', '')
|
||||
if req_id:
|
||||
await self.bot.reply_text(req_id, fallback)
|
||||
else:
|
||||
await self.bot.set_message(msg_id, fallback)
|
||||
return {'stream': False, 'form': True, 'fallback': True}
|
||||
|
||||
if _ws_mode:
|
||||
success = await self.bot.push_stream_chunk(msg_id, content, is_final=is_final)
|
||||
if not success and is_final:
|
||||
@@ -480,142 +385,6 @@ class WecomBotAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
"""Whether streaming output is enabled for this bot instance."""
|
||||
return self.config.get('enable-stream-reply', True)
|
||||
|
||||
async def _handle_synthetic_chunk(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
bot_message,
|
||||
content: str,
|
||||
is_final: bool,
|
||||
ws_mode: bool,
|
||||
) -> dict:
|
||||
"""Handle reply_message_chunk for synthetic events (button clicks).
|
||||
|
||||
Synthetic events have no inbound message → no msg_id, no req_id,
|
||||
no stream session. We can't do incremental streaming, so we
|
||||
buffer chunks per-conversation and flush on ``is_final`` via the
|
||||
proactive send path.
|
||||
|
||||
Buffer keyed by ``(launcher_type, launcher_id)`` from the
|
||||
synthetic event itself. Only ws mode has a usable proactive-send
|
||||
path right now (``ws_client.send_message`` /
|
||||
``ws_client.send_template_card``); webhook mode requires a
|
||||
corpid/secret we don't have, so it logs and drops.
|
||||
"""
|
||||
if isinstance(message_source, platform_events.GroupMessage):
|
||||
chat_id = str(message_source.group.id)
|
||||
else:
|
||||
chat_id = str(message_source.sender.id)
|
||||
|
||||
form_data = getattr(bot_message, '_form_data', None)
|
||||
|
||||
# Buffer streaming content until is_final.
|
||||
buf_key = chat_id
|
||||
if not hasattr(self, '_synthetic_buffers'):
|
||||
# Attribute-not-declared trick: pydantic forbids dynamic attrs
|
||||
# on the model, but plain instance dicts via object.__setattr__
|
||||
# do work. Lazy-create on first call.
|
||||
object.__setattr__(self, '_synthetic_buffers', {})
|
||||
buffers: dict[str, str] = self._synthetic_buffers
|
||||
if content and not form_data:
|
||||
previous = buffers.get(buf_key, '')
|
||||
if previous and content.startswith(previous):
|
||||
buffers[buf_key] = content
|
||||
elif previous and previous.endswith(content):
|
||||
buffers[buf_key] = previous
|
||||
else:
|
||||
buffers[buf_key] = previous + content
|
||||
|
||||
if not is_final:
|
||||
return {'stream': True, 'synthetic': True, 'buffered': True}
|
||||
|
||||
final_content = buffers.pop(buf_key, '')
|
||||
if content:
|
||||
if final_content and content.startswith(final_content):
|
||||
final_content = content
|
||||
elif final_content and final_content.endswith(content):
|
||||
pass
|
||||
else:
|
||||
final_content = final_content + content
|
||||
|
||||
if not ws_mode:
|
||||
await self.logger.warning(
|
||||
'WeComBot webhook mode cannot proactively push synthetic-event '
|
||||
'output (no corpid/secret); the resume reply is dropped. '
|
||||
f'content_len={len(final_content)} form_data_present={form_data is not None}'
|
||||
)
|
||||
return {'stream': False, 'synthetic': True, 'dropped': True}
|
||||
|
||||
# ws mode: proactive send.
|
||||
try:
|
||||
if form_data:
|
||||
# Determine user_id / chat_id for the routing context of any
|
||||
# subsequent click on this card.
|
||||
if isinstance(message_source, platform_events.GroupMessage):
|
||||
routing_chat_id = str(message_source.group.id)
|
||||
routing_user_id = str(message_source.sender.id)
|
||||
else:
|
||||
routing_chat_id = ''
|
||||
routing_user_id = str(message_source.sender.id)
|
||||
payload = self._build_button_interaction_payload_from_form(
|
||||
form_data,
|
||||
user_id=routing_user_id,
|
||||
chat_id=routing_chat_id,
|
||||
)
|
||||
await self.bot.send_template_card(chat_id, payload)
|
||||
await self.logger.info(
|
||||
f'WeComBot ws: proactively sent template_card for synthetic event '
|
||||
f'chat_id={chat_id} form_token={form_data.get("form_token")!r} '
|
||||
f'workflow_run_id={form_data.get("workflow_run_id")!r}'
|
||||
)
|
||||
elif final_content:
|
||||
await self.bot.send_message(chat_id, final_content)
|
||||
await self.logger.info(
|
||||
f'WeComBot ws: proactively sent text for synthetic event chat_id={chat_id} len={len(final_content)}'
|
||||
)
|
||||
except Exception:
|
||||
await self.logger.error(f'WeComBot: synthetic event proactive send failed: {traceback.format_exc()}')
|
||||
return {'stream': False, 'synthetic': True, 'error': True}
|
||||
|
||||
return {'stream': True, 'synthetic': True}
|
||||
|
||||
def _build_button_interaction_payload_from_form(
|
||||
self, form_data: dict, *, user_id: str = '', chat_id: str = ''
|
||||
) -> dict:
|
||||
"""Build a button_interaction payload + track task_id for click resolution.
|
||||
|
||||
Unlike the inbound-event path (where push_form_pause registers the
|
||||
task_id with the active stream session), proactive sends still
|
||||
need the task_id registered so button clicks find pending_form.
|
||||
For ws mode we stash it directly on the ws_client's pending dict.
|
||||
"""
|
||||
from langbot.libs.wecom_ai_bot_api.api import build_human_input_template_card_payload
|
||||
import secrets as _secrets
|
||||
|
||||
task_id = f'dify-{_secrets.token_hex(12)}'
|
||||
source = getattr(self.bot, 'card_source', None)
|
||||
payload = build_human_input_template_card_payload(
|
||||
form_data,
|
||||
task_id,
|
||||
source=source,
|
||||
select_as_buttons=not self.config.get('enable-webhook', False),
|
||||
)
|
||||
|
||||
# Register task_id → form_data so the click callback can find it.
|
||||
# user_id / chat_id are required so _on_card_action can route the
|
||||
# resulting synthetic query back to the right user. msg_id / req_id
|
||||
# / stream_id are intentionally empty — synthetic cards have no
|
||||
# inbound message to anchor on.
|
||||
if hasattr(self.bot, '_pending_forms_by_task'):
|
||||
self.bot._pending_forms_by_task[task_id] = {
|
||||
'form_data': form_data,
|
||||
'msg_id': '',
|
||||
'user_id': user_id,
|
||||
'chat_id': chat_id,
|
||||
'stream_id': '',
|
||||
'req_id': '',
|
||||
}
|
||||
return payload
|
||||
|
||||
async def send_message(self, target_type, target_id, message):
|
||||
_ws_mode = not self.config.get('enable-webhook', False)
|
||||
if _ws_mode:
|
||||
@@ -762,191 +531,3 @@ class WecomBotAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
|
||||
async def is_muted(self, group_id: int) -> bool:
|
||||
pass
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Dify human-input button-interaction click handling
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _on_card_action(self, session, action_id: str, task_id: str, raw_event: dict) -> None:
|
||||
"""Translate a button click on a button_interaction card into a
|
||||
synthetic ``_dify_form_action`` query enqueued on the pool.
|
||||
|
||||
Pattern mirrors DingTalk / Lark / Telegram so the runner's
|
||||
``_merge_pending_form_action`` path resumes the workflow.
|
||||
"""
|
||||
import langbot_plugin.api.entities.builtin.provider.session as provider_session
|
||||
|
||||
form = session.pending_form or {}
|
||||
await self.logger.info(
|
||||
f'WeComBot _on_card_action: task_id={task_id} action_id={action_id!r} '
|
||||
f'form_token={form.get("form_token")!r} workflow_run_id={form.get("workflow_run_id")!r} '
|
||||
f'session.user_id={session.user_id!r} session.chat_id={session.chat_id!r}'
|
||||
)
|
||||
|
||||
actions = form.get('actions') or []
|
||||
tce = extract_template_card_event_payload(raw_event) if isinstance(raw_event, dict) else {}
|
||||
_, _, card_type = extract_template_card_action(tce)
|
||||
selections = extract_template_card_selections(tce, form)
|
||||
if not selections:
|
||||
selections = parse_select_button_action(action_id, form)
|
||||
await self.logger.info(
|
||||
f'WeComBot template_card selections: task_id={task_id} card_type={card_type} selections={selections}'
|
||||
)
|
||||
if card_type == 'multiple_interaction' and not selections:
|
||||
await self.logger.warning(
|
||||
f'WeComBot: multiple_interaction callback has no parseable selections; raw={str(tce)[:1000]}'
|
||||
)
|
||||
return
|
||||
is_select_submit = card_type == 'multiple_interaction' or bool(selections)
|
||||
|
||||
clean_action_id = '' if is_select_submit else (action_id or '').strip()
|
||||
action_title = clean_action_id
|
||||
for a in actions:
|
||||
if str(a.get('id', '')) == clean_action_id:
|
||||
action_title = a.get('title') or clean_action_id
|
||||
break
|
||||
|
||||
inputs = dict(form.get('inputs') or {})
|
||||
inputs.update(selections)
|
||||
|
||||
def _missing_fields_after_select() -> list[str]:
|
||||
missing: list[str] = []
|
||||
for field in form.get('input_defs') or form.get('all_input_defs') or []:
|
||||
field_name = str(field.get('output_variable_name') or '').strip()
|
||||
if not field_name:
|
||||
continue
|
||||
if inputs.get(field_name) in (None, '', []):
|
||||
missing.append(field_name)
|
||||
return missing
|
||||
|
||||
input_progress = False
|
||||
if is_select_submit:
|
||||
missing_fields = _missing_fields_after_select()
|
||||
if not missing_fields and len(actions) == 1:
|
||||
action = actions[0]
|
||||
clean_action_id = str(action.get('id') or '').strip()
|
||||
action_title = action.get('title') or clean_action_id
|
||||
elif not missing_fields and len(actions) > 1:
|
||||
if not self.config.get('enable-webhook', False):
|
||||
action_form_data = {
|
||||
'form_content': form.get('raw_form_content') or form.get('form_content') or '',
|
||||
'raw_form_content': form.get('raw_form_content') or form.get('form_content') or '',
|
||||
'input_defs': [],
|
||||
'all_input_defs': form.get('all_input_defs') or form.get('input_defs') or [],
|
||||
'inputs': inputs,
|
||||
'actions': actions,
|
||||
'node_title': form.get('node_title', ''),
|
||||
'workflow_run_id': form.get('workflow_run_id', ''),
|
||||
'form_token': form.get('form_token', ''),
|
||||
'pipeline_uuid': form.get('pipeline_uuid', ''),
|
||||
'_action_select_only': True,
|
||||
}
|
||||
target_chat_id = session.chat_id or session.user_id or ''
|
||||
try:
|
||||
payload = self._build_button_interaction_payload_from_form(
|
||||
action_form_data,
|
||||
user_id=session.user_id or '',
|
||||
chat_id=session.chat_id or '',
|
||||
)
|
||||
await self.bot.send_template_card(target_chat_id, payload)
|
||||
await self.logger.info(
|
||||
f'WeComBot: sent action-select button card after select submit '
|
||||
f'task_id={task_id} action_count={len(actions)}'
|
||||
)
|
||||
except Exception:
|
||||
await self.logger.error(
|
||||
f'WeComBot: failed to send action-select button card: {traceback.format_exc()}'
|
||||
)
|
||||
return
|
||||
await self.logger.warning(
|
||||
'WeComBot webhook mode cannot proactively send action-select button card after select submit'
|
||||
)
|
||||
return
|
||||
else:
|
||||
input_progress = True
|
||||
action_title = 'Submit'
|
||||
|
||||
launcher_id = session.user_id or session.chat_id or ''
|
||||
sender_user_id = session.user_id or launcher_id
|
||||
# WeCom AI bot has both single-chat and group-chat; chat_id present
|
||||
# indicates group context.
|
||||
if session.chat_id:
|
||||
launcher_type = provider_session.LauncherTypes.GROUP
|
||||
launcher_id = session.chat_id
|
||||
else:
|
||||
launcher_type = provider_session.LauncherTypes.PERSON
|
||||
launcher_id = session.user_id or ''
|
||||
|
||||
form_action_data = {
|
||||
'form_token': form.get('form_token', ''),
|
||||
'workflow_run_id': form.get('workflow_run_id', ''),
|
||||
'action_id': clean_action_id,
|
||||
'action_title': action_title,
|
||||
'node_title': form.get('node_title', ''),
|
||||
'user': f'{launcher_type.value}_{launcher_id}',
|
||||
'inputs': inputs,
|
||||
}
|
||||
if input_progress:
|
||||
form_action_data['_input_progress'] = True
|
||||
|
||||
message_chain = platform_message.MessageChain([platform_message.Plain(text=f'[Form Action: {action_title}]')])
|
||||
|
||||
if launcher_type == provider_session.LauncherTypes.GROUP:
|
||||
synthetic_event = platform_events.GroupMessage(
|
||||
sender=platform_entities.GroupMember(
|
||||
id=sender_user_id,
|
||||
member_name='',
|
||||
permission=platform_entities.Permission.Member,
|
||||
group=platform_entities.Group(
|
||||
id=launcher_id,
|
||||
name='',
|
||||
permission=platform_entities.Permission.Member,
|
||||
),
|
||||
special_title='',
|
||||
),
|
||||
message_chain=message_chain,
|
||||
time=int(time.time()),
|
||||
source_platform_object=None,
|
||||
)
|
||||
else:
|
||||
synthetic_event = platform_events.FriendMessage(
|
||||
sender=platform_entities.Friend(
|
||||
id=sender_user_id,
|
||||
nickname='',
|
||||
remark='',
|
||||
),
|
||||
message_chain=message_chain,
|
||||
time=int(time.time()),
|
||||
source_platform_object=None,
|
||||
)
|
||||
|
||||
if self.ap is None:
|
||||
await self.logger.error('WeComBot: ap not injected; cannot enqueue button-click query')
|
||||
return
|
||||
|
||||
bot_uuid = ''
|
||||
pipeline_uuid = form.get('pipeline_uuid') or None
|
||||
for bot in self.ap.platform_mgr.bots:
|
||||
if bot.adapter is self:
|
||||
bot_uuid = bot.bot_entity.uuid
|
||||
pipeline_uuid = pipeline_uuid or bot.bot_entity.use_pipeline_uuid
|
||||
break
|
||||
|
||||
try:
|
||||
await self.ap.query_pool.add_query(
|
||||
bot_uuid=bot_uuid,
|
||||
launcher_type=launcher_type,
|
||||
launcher_id=launcher_id,
|
||||
sender_id=sender_user_id,
|
||||
message_event=synthetic_event,
|
||||
message_chain=message_chain,
|
||||
adapter=self,
|
||||
pipeline_uuid=pipeline_uuid,
|
||||
variables={
|
||||
'_dify_form_action': form_action_data,
|
||||
'_routed_by_rule': True,
|
||||
},
|
||||
)
|
||||
await self.logger.info(f'WeComBot: button-click query enqueued action_id={clean_action_id!r}')
|
||||
except Exception:
|
||||
await self.logger.error(f'WeComBot: enqueue button-click query failed: {traceback.format_exc()}')
|
||||
|
||||
@@ -2,7 +2,6 @@ from __future__ import annotations
|
||||
import typing
|
||||
import asyncio
|
||||
import traceback
|
||||
import uuid
|
||||
|
||||
import datetime
|
||||
import pydantic
|
||||
@@ -183,28 +182,7 @@ class WecomCSAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
)
|
||||
|
||||
async def send_message(self, target_type: str, target_id: str, message: platform_message.MessageChain):
|
||||
if target_type != 'person':
|
||||
raise ValueError('WeCom customer service only supports sending messages to person targets')
|
||||
|
||||
open_kfid = self.bot_account_id
|
||||
external_userid = target_id
|
||||
if '|' in target_id:
|
||||
open_kfid, external_userid = target_id.split('|', 1)
|
||||
if external_userid.startswith('u'):
|
||||
external_userid = external_userid[1:]
|
||||
if not open_kfid:
|
||||
raise ValueError('WeCom customer service open_kfid is required before sending messages')
|
||||
|
||||
content_list = await WecomMessageConverter.yiri2target(message, self.bot)
|
||||
for content in content_list:
|
||||
msgid = f'langbot_{uuid.uuid4().hex}'
|
||||
if content['type'] == 'text':
|
||||
await self.bot.send_text_msg(
|
||||
open_kfid=open_kfid,
|
||||
external_userid=external_userid,
|
||||
msgid=msgid,
|
||||
content=content['content'],
|
||||
)
|
||||
pass
|
||||
|
||||
def set_bot_uuid(self, bot_uuid: str):
|
||||
"""设置 bot UUID(用于生成 webhook URL)"""
|
||||
|
||||
@@ -363,9 +363,19 @@ class RuntimeConnectionHandler(handler.Handler):
|
||||
extra_args=extra_args,
|
||||
)
|
||||
|
||||
# invoke_llm returns (message, usage_info) tuple
|
||||
if isinstance(result, tuple) and len(result) == 2:
|
||||
msg, usage_info = result
|
||||
msg_dump = msg.model_dump()
|
||||
# Attach usage info to message dump
|
||||
if usage_info:
|
||||
msg_dump['usage'] = usage_info
|
||||
else:
|
||||
msg_dump = result.model_dump()
|
||||
|
||||
return handler.ActionResponse.success(
|
||||
data={
|
||||
'message': result.model_dump(),
|
||||
'message': msg_dump,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -81,8 +81,13 @@ class RuntimeProvider:
|
||||
msg, usage_info = result
|
||||
if usage_info:
|
||||
_store_llm_usage(query, usage_info)
|
||||
input_tokens = usage_info.get('prompt_tokens', 0)
|
||||
output_tokens = usage_info.get('completion_tokens', 0)
|
||||
input_tokens = usage_info.get('prompt_tokens', usage_info.get('input_tokens', 0))
|
||||
output_tokens = usage_info.get('completion_tokens', usage_info.get('output_tokens', 0))
|
||||
# Attach usage info to message using object.__setattr__ to bypass pydantic validation
|
||||
try:
|
||||
object.__setattr__(msg, 'usage', usage_info)
|
||||
except (AttributeError, TypeError):
|
||||
pass # If we can't set it, just skip it
|
||||
return msg
|
||||
else:
|
||||
return result
|
||||
@@ -98,7 +103,7 @@ class RuntimeProvider:
|
||||
|
||||
# Import monitoring helper
|
||||
try:
|
||||
from ...pipeline import monitoring_helper
|
||||
from ...pipeline import monitor
|
||||
|
||||
# Get monitoring metadata from query variables
|
||||
if query.variables:
|
||||
@@ -110,7 +115,7 @@ class RuntimeProvider:
|
||||
pipeline_name = 'Unknown'
|
||||
message_id = None
|
||||
|
||||
await monitoring_helper.MonitoringHelper.record_llm_call(
|
||||
await monitor.MonitoringHelper.record_llm_call(
|
||||
ap=self.requester.ap,
|
||||
query=query,
|
||||
bot_id=query.bot_uuid or 'unknown',
|
||||
@@ -177,7 +182,7 @@ class RuntimeProvider:
|
||||
|
||||
# Import monitoring helper
|
||||
try:
|
||||
from ...pipeline import monitoring_helper
|
||||
from ...pipeline import monitor
|
||||
|
||||
# Get monitoring metadata from query variables
|
||||
if query.variables:
|
||||
@@ -189,7 +194,7 @@ class RuntimeProvider:
|
||||
pipeline_name = 'Unknown'
|
||||
message_id = None
|
||||
|
||||
await monitoring_helper.MonitoringHelper.record_llm_call(
|
||||
await monitor.MonitoringHelper.record_llm_call(
|
||||
ap=self.requester.ap,
|
||||
query=query,
|
||||
bot_id=query.bot_uuid or 'unknown',
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -6,7 +6,7 @@ import json
|
||||
import re
|
||||
import time
|
||||
import typing
|
||||
from contextlib import AsyncExitStack, asynccontextmanager
|
||||
from contextlib import AsyncExitStack
|
||||
import traceback
|
||||
from langbot_plugin.api.entities.events import pipeline_query
|
||||
import sqlalchemy
|
||||
@@ -18,7 +18,6 @@ from mcp import ClientSession, StdioServerParameters, types as mcp_types
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from mcp.shared.exceptions import McpError
|
||||
from pydantic import AnyUrl
|
||||
|
||||
from .. import loader
|
||||
@@ -26,7 +25,7 @@ from ....core import app
|
||||
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
|
||||
import langbot_plugin.api.entities.builtin.provider.message as provider_message
|
||||
from ....entity.persistence import mcp as persistence_mcp
|
||||
from .mcp_stdio import BoxStdioSessionRuntime, MCPServerBoxConfig, MCPSessionErrorPhase, _ColdStartRetry # noqa: F401
|
||||
from .mcp_stdio import BoxStdioSessionRuntime, MCPServerBoxConfig, MCPSessionErrorPhase # noqa: F401
|
||||
|
||||
# Synthesized LLM tools for MCP resources (not from server tools/list).
|
||||
# Dispatched in MCPLoader.invoke_tool; placeholder func on LLMTool is never used.
|
||||
@@ -186,16 +185,6 @@ class MCPSessionStatus(enum.Enum):
|
||||
ERROR = 'error'
|
||||
|
||||
|
||||
class _TransportReconnect(Exception):
|
||||
"""Internal signal: the Box stdio WS transport dropped but the managed
|
||||
process is still alive. Triggers a lightweight transport reconnect that
|
||||
reuses the live process, instead of a full process rebuild.
|
||||
|
||||
Reconnect attempts are NOT counted toward the fatal retry budget, so a
|
||||
long-lived session can survive arbitrarily many transient drops.
|
||||
"""
|
||||
|
||||
|
||||
class RuntimeMCPSession:
|
||||
"""运行时 MCP 会话"""
|
||||
|
||||
@@ -265,16 +254,6 @@ class RuntimeMCPSession:
|
||||
self._lifecycle_task = None
|
||||
self._shutdown_event = asyncio.Event()
|
||||
self._ready_event = asyncio.Event()
|
||||
# Set transiently when a WS transport drop should NOT stop the managed
|
||||
# process (it will be re-attached on the next initialize()).
|
||||
self._preserve_managed_process = False
|
||||
|
||||
# Log buffer for capturing stderr from Box managed process (maxlen=500 keeps
|
||||
# recent lines without unbounded memory growth)
|
||||
import collections as _collections
|
||||
|
||||
self._log_buffer: _collections.deque = _collections.deque(maxlen=500)
|
||||
self._last_stderr_text: str = ''
|
||||
|
||||
self._box_stdio_runtime = BoxStdioSessionRuntime(self)
|
||||
self.box_config = self._box_stdio_runtime.config
|
||||
@@ -336,34 +315,23 @@ class RuntimeMCPSession:
|
||||
|
||||
await self.session.initialize()
|
||||
|
||||
@asynccontextmanager
|
||||
async def _streamable_http_session(self) -> typing.AsyncIterator[ClientSession]:
|
||||
"""Enter a fully initialized Streamable HTTP session as one context.
|
||||
|
||||
Initialization must happen inside the same context manager that owns the
|
||||
MCP transport. The SDK reports request failures by cancelling the host
|
||||
task and raises the real HTTP error from its TaskGroup during context
|
||||
exit. Keeping these nested contexts together guarantees a failed
|
||||
``__aenter__`` unwinds immediately, so callers see the HTTPStatusError
|
||||
instead of a detached CancelledError. It also owns the injected HTTPX
|
||||
client, which the MCP SDK deliberately does not close for callers.
|
||||
"""
|
||||
async with httpx.AsyncClient(
|
||||
headers=self.server_config.get('headers', {}),
|
||||
timeout=self.server_config.get('timeout', 10),
|
||||
follow_redirects=True,
|
||||
) as http_client:
|
||||
async with streamable_http_client(
|
||||
self.server_config['url'],
|
||||
http_client=http_client,
|
||||
) as transport:
|
||||
read, write, _ = transport
|
||||
async with ClientSession(read, write) as session:
|
||||
await session.initialize()
|
||||
yield session
|
||||
|
||||
async def _init_streamable_http_server(self):
|
||||
self.session = await self.exit_stack.enter_async_context(self._streamable_http_session())
|
||||
transport = await self.exit_stack.enter_async_context(
|
||||
streamable_http_client(
|
||||
self.server_config['url'],
|
||||
http_client=httpx.AsyncClient(
|
||||
headers=self.server_config.get('headers', {}),
|
||||
timeout=self.server_config.get('timeout', 10),
|
||||
follow_redirects=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
read, write, _ = transport
|
||||
|
||||
self.session = await self.exit_stack.enter_async_context(ClientSession(read, write))
|
||||
|
||||
await self.session.initialize()
|
||||
|
||||
async def _init_remote_server(self):
|
||||
"""Connect to a remote MCP server, auto-detecting the transport.
|
||||
@@ -378,15 +346,9 @@ class RuntimeMCPSession:
|
||||
await self._init_streamable_http_server()
|
||||
return
|
||||
except Exception as e:
|
||||
if not self._should_fallback_to_sse(e):
|
||||
self.ap.logger.info(
|
||||
f'MCP server {self.server_name}: Streamable HTTP transport failed '
|
||||
f'({self._describe_exception(e)}); not falling back to SSE'
|
||||
)
|
||||
raise
|
||||
self.ap.logger.info(
|
||||
f'MCP server {self.server_name}: Streamable HTTP initialize failed with a compatible HTTP status '
|
||||
f'({self._describe_exception(e)}), falling back to legacy SSE'
|
||||
f'MCP server {self.server_name}: Streamable HTTP transport failed '
|
||||
f'({self._describe_exception(e)}), falling back to SSE'
|
||||
)
|
||||
|
||||
# The Streamable HTTP attempt may have partially entered the transport /
|
||||
@@ -437,39 +399,11 @@ class RuntimeMCPSession:
|
||||
task.cancel()
|
||||
for task in done:
|
||||
if task is monitor_task and not self._shutdown_event.is_set():
|
||||
# The monitor completed. This is EITHER the managed
|
||||
# process actually exiting OR just the WS transport
|
||||
# dropping while the process stays alive in the Box
|
||||
# runtime. Re-check the real process state so a
|
||||
# transient transport drop reconnects (reusing the live
|
||||
# process) instead of tearing the process down and
|
||||
# running a full rebuild+backoff cycle.
|
||||
process_still_running = False
|
||||
try:
|
||||
process_still_running = await self._box_stdio_runtime._managed_process_is_running()
|
||||
except Exception:
|
||||
process_still_running = False
|
||||
if process_still_running:
|
||||
self.ap.logger.info(
|
||||
f'MCP server {self.server_name}: transport dropped but '
|
||||
f'managed process is still running; reconnecting transport'
|
||||
)
|
||||
self.error_phase = MCPSessionErrorPhase.RELAY_CONNECT
|
||||
# Preserve the live process across the finally-block
|
||||
# cleanup: only the WS transport should be torn down.
|
||||
self._preserve_managed_process = True
|
||||
raise _TransportReconnect('Box managed process transport dropped; reconnecting')
|
||||
self.error_phase = MCPSessionErrorPhase.RUNTIME
|
||||
raise Exception('Box managed process exited unexpectedly')
|
||||
else:
|
||||
await self._shutdown_event.wait()
|
||||
|
||||
except _ColdStartRetry:
|
||||
# Cold-start in progress: set the preserve flag BEFORE the finally
|
||||
# block runs so it does not stop the live managed process. The outer
|
||||
# _lifecycle_loop_with_retry will reuse it on the next attempt.
|
||||
self._preserve_managed_process = True
|
||||
raise
|
||||
except Exception as e:
|
||||
self.status = MCPSessionStatus.ERROR
|
||||
self.error_message = str(e)
|
||||
@@ -490,55 +424,14 @@ class RuntimeMCPSession:
|
||||
except Exception as e:
|
||||
self.ap.logger.error(f'Error cleaning up MCP session {self.server_name}: {e}\n{traceback.format_exc()}')
|
||||
finally:
|
||||
# On a transport-only reconnect the managed process is healthy
|
||||
# and will be re-attached on the next initialize(); do NOT stop
|
||||
# it. Any other exit path fully tears the session down.
|
||||
if getattr(self, '_preserve_managed_process', False):
|
||||
self._preserve_managed_process = False
|
||||
else:
|
||||
await self._cleanup_box_stdio_session()
|
||||
await self._cleanup_box_stdio_session()
|
||||
|
||||
async def _lifecycle_loop_with_retry(self):
|
||||
"""Wrap _lifecycle_loop with retry and exponential backoff."""
|
||||
attempt = 0
|
||||
while attempt <= self._MAX_RETRIES:
|
||||
for attempt in range(self._MAX_RETRIES + 1):
|
||||
try:
|
||||
await self._lifecycle_loop()
|
||||
return # Normal shutdown, don't retry
|
||||
except _TransportReconnect as e:
|
||||
# Transient WS transport drop while the managed process is still
|
||||
# alive. Reconnect promptly WITHOUT consuming the fatal retry
|
||||
# budget and WITHOUT stopping the process — initialize() will
|
||||
# re-attach to the live process. This is what lets a long-lived
|
||||
# stdio MCP survive repeated brief event-loop stalls / pings.
|
||||
if self._shutdown_event.is_set():
|
||||
return
|
||||
self.ap.logger.info(
|
||||
f'MCP session {self.server_name}: reconnecting transport ({self._describe_exception(e)})'
|
||||
)
|
||||
self.status = MCPSessionStatus.CONNECTING
|
||||
self.error_message = None
|
||||
self.error_phase = None
|
||||
await asyncio.sleep(1)
|
||||
continue
|
||||
except _ColdStartRetry as e:
|
||||
# The managed process is alive but still cold-starting (e.g.
|
||||
# `npx -y <pkg>` is still installing) and cannot yet answer the
|
||||
# handshake. Reuse the live process and retry the attach WITHOUT
|
||||
# consuming the fatal retry budget or stopping the process, so a
|
||||
# slow cold start is waited out instead of failing. Preserve the
|
||||
# process across the finally-block cleanup.
|
||||
if self._shutdown_event.is_set():
|
||||
return
|
||||
self._preserve_managed_process = True
|
||||
self.ap.logger.debug(
|
||||
f'MCP session {self.server_name}: waiting for cold start ({self._describe_exception(e)})'
|
||||
)
|
||||
self.status = MCPSessionStatus.CONNECTING
|
||||
self.error_message = None
|
||||
self.error_phase = None
|
||||
await asyncio.sleep(2)
|
||||
continue
|
||||
except Exception as e:
|
||||
self.retry_count = attempt + 1
|
||||
if self._shutdown_event.is_set():
|
||||
@@ -567,7 +460,6 @@ class RuntimeMCPSession:
|
||||
self.error_message = None
|
||||
self.error_phase = None
|
||||
await asyncio.sleep(delay)
|
||||
attempt += 1
|
||||
|
||||
@staticmethod
|
||||
def _describe_exception(exc: BaseException) -> str:
|
||||
@@ -593,36 +485,6 @@ class RuntimeMCPSession:
|
||||
unique = [m for m in leaves if not (m in seen or seen.add(m))]
|
||||
return '; '.join(unique) if unique else f'{type(exc).__name__}: {exc}'
|
||||
|
||||
@staticmethod
|
||||
def _iter_exception_leaves(exc: BaseException) -> typing.Iterator[BaseException]:
|
||||
sub = getattr(exc, 'exceptions', None)
|
||||
if sub: # ExceptionGroup / BaseExceptionGroup
|
||||
for child in sub:
|
||||
yield from RuntimeMCPSession._iter_exception_leaves(child)
|
||||
else:
|
||||
yield exc
|
||||
|
||||
@staticmethod
|
||||
def _should_fallback_to_sse(exc: BaseException) -> bool:
|
||||
"""Whether a Streamable HTTP failure matches legacy-SSE fallback.
|
||||
|
||||
Only protocol-compatibility responses trigger fallback. Authentication,
|
||||
authorization, throttling, and server failures must remain visible
|
||||
instead of being retried against a different transport.
|
||||
|
||||
MCP SDK 1.26 translates an HTTP 404 initialize response into a synthetic
|
||||
``McpError(32600, 'Session terminated')`` rather than preserving the
|
||||
HTTPStatusError, so recognize that exact SDK sentinel as 404-compatible.
|
||||
"""
|
||||
fallback_statuses = {400, 404, 405}
|
||||
for leaf in RuntimeMCPSession._iter_exception_leaves(exc):
|
||||
if isinstance(leaf, httpx.HTTPStatusError):
|
||||
if leaf.response.status_code in fallback_statuses:
|
||||
return True
|
||||
elif isinstance(leaf, McpError) and leaf.error.code == 32600 and leaf.error.message == 'Session terminated':
|
||||
return True
|
||||
return False
|
||||
|
||||
_MONITOR_POLL_INTERVAL = 5
|
||||
_MONITOR_MAX_CONSECUTIVE_ERRORS = 3
|
||||
|
||||
@@ -1065,14 +927,11 @@ class RuntimeMCPSession:
|
||||
return self._box_stdio_runtime.uses_box_stdio()
|
||||
|
||||
def _build_box_session_id(self) -> str:
|
||||
# Both live servers and transient config-page tests share ONE Box
|
||||
# session ('mcp-shared'). A test therefore reuses the already-running
|
||||
# container (and, for an existing server, its live managed process)
|
||||
# instead of paying a full per-test session cold-start + dependency
|
||||
# bootstrap. Isolation between a test and the live servers is provided
|
||||
# at the *process* level: each server/test has its own process_id and a
|
||||
# test only ever stops its own process_id (see cleanup_session), so it
|
||||
# never disturbs another server's process or the shared session itself.
|
||||
# Transient test sessions get their own isolated Box session so a
|
||||
# failing/short-lived test can never disturb the shared session that
|
||||
# hosts live, already-connected MCP servers.
|
||||
if self.is_transient:
|
||||
return f'mcp-test-{self.server_uuid}'
|
||||
return 'mcp-shared'
|
||||
|
||||
def _rewrite_path(self, path: str, host_path: str | None) -> str:
|
||||
|
||||
@@ -6,7 +6,7 @@ import os
|
||||
import shutil
|
||||
import shlex
|
||||
import threading
|
||||
from contextlib import suppress, AsyncExitStack
|
||||
from contextlib import suppress
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pydantic
|
||||
@@ -57,23 +57,6 @@ class MCPSessionErrorPhase(enum.Enum):
|
||||
BOX_UNAVAILABLE = 'box_unavailable'
|
||||
|
||||
|
||||
def _get_default_memory_mb(ap) -> int:
|
||||
"""Read box.default_memory_mb from instance config (env: BOX__DEFAULT_MEMORY_MB).
|
||||
|
||||
Falls back to 1536 MB — a safe floor for Node.js V8 + WASM under nsjail.
|
||||
Operators running memory-constrained hosts can lower this; those with large
|
||||
machines can raise it. Individual MCP servers can still override via their
|
||||
own box.memory_mb setting.
|
||||
"""
|
||||
try:
|
||||
data = getattr(getattr(ap, 'instance_config', None), 'data', None)
|
||||
if isinstance(data, dict):
|
||||
return int(data.get('box', {}).get('default_memory_mb', 1536))
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
return 1536
|
||||
|
||||
|
||||
class MCPServerBoxConfig(pydantic.BaseModel):
|
||||
"""Structured configuration for running an MCP server inside a Box container."""
|
||||
|
||||
@@ -91,35 +74,6 @@ class MCPServerBoxConfig(pydantic.BaseModel):
|
||||
model_config = pydantic.ConfigDict(extra='ignore')
|
||||
|
||||
|
||||
_HANDSHAKE_ATTEMPT_TIMEOUT_SEC = 10.0
|
||||
|
||||
|
||||
class _TransferredStack:
|
||||
"""Adapts an already-populated AsyncExitStack into an async context manager
|
||||
so ownership of its resources can be transferred into another exit stack.
|
||||
Entering is a no-op; exiting closes the wrapped stack (and thus the live WS
|
||||
transport + ClientSession) when the owning session shuts down."""
|
||||
|
||||
def __init__(self, stack: AsyncExitStack):
|
||||
self._stack = stack
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
await self._stack.aclose()
|
||||
return False
|
||||
|
||||
|
||||
class _ColdStartRetry(Exception):
|
||||
"""Signal: the managed process is alive but not yet answering the MCP
|
||||
handshake because it is still cold-starting (e.g. `npx -y <pkg>` is still
|
||||
installing). The outer lifecycle retry treats this like a transient
|
||||
reconnect: it reuses the live process and does not count toward the fatal
|
||||
retry budget, so a slow cold start is waited out rather than failing.
|
||||
"""
|
||||
|
||||
|
||||
class BoxStdioSessionRuntime:
|
||||
"""Encapsulate Box-backed stdio MCP session orchestration."""
|
||||
|
||||
@@ -159,17 +113,7 @@ class BoxStdioSessionRuntime:
|
||||
read_only_rootfs=self.config.read_only_rootfs if self.config.read_only_rootfs is not None else False,
|
||||
image=self.config.image,
|
||||
cpus=self.config.cpus,
|
||||
# Node.js runtimes (npx/bunx) reserve large virtual address space and
|
||||
# load WebAssembly modules (llhttp) on startup; the default 512 MB
|
||||
# cgroup_mem_max is too small and causes OOM kills (return_code=137).
|
||||
# Auto-bump to 1024 MB when the runner is npx/bunx/pnpm dlx.
|
||||
# Per-server override wins; global default comes from
|
||||
# config.yaml box.default_memory_mb (env: BOX__DEFAULT_MEMORY_MB).
|
||||
# Hard floor of 1536 MB: enough for Node.js V8 + WASM without OOM.
|
||||
# Per-server override wins; global default from config.yaml
|
||||
# box.default_memory_mb (env: BOX__DEFAULT_MEMORY_MB), hard floor
|
||||
# of 1536 MB so Node.js V8 + WASM never OOM under nsjail.
|
||||
memory_mb=(self.config.memory_mb or _get_default_memory_mb(self.ap)),
|
||||
memory_mb=self.config.memory_mb,
|
||||
pids_limit=self.config.pids_limit,
|
||||
persistent=True,
|
||||
)
|
||||
@@ -229,55 +173,28 @@ class BoxStdioSessionRuntime:
|
||||
stderr_preview = (result.stderr or '')[:500]
|
||||
raise Exception(f'Dependency install failed (exit code {result.exit_code}): {stderr_preview}')
|
||||
|
||||
# Reuse an already-running managed process instead of rebuilding it.
|
||||
# The Box runtime keeps the managed process alive across a transient
|
||||
# WebSocket transport drop, so on a reconnect we only need to re-attach
|
||||
# the WS below. Rebuilding here would needlessly stop a healthy process
|
||||
# and re-run the (slow, network-touching) dependency bootstrap.
|
||||
if not await self._managed_process_is_running():
|
||||
try:
|
||||
process_workspace = (
|
||||
self._build_workspace(host_path=host_path, workdir=process_cwd, mount_path=process_cwd)
|
||||
if host_path
|
||||
else workspace
|
||||
)
|
||||
payload = process_workspace.build_process_payload(
|
||||
self.server_config['command'],
|
||||
self.server_config.get('args', []),
|
||||
env=self.server_config.get('env', {}),
|
||||
cwd=process_cwd,
|
||||
)
|
||||
if install_cmd:
|
||||
payload = self._wrap_process_payload_with_python_env(payload, process_cwd)
|
||||
payload['process_id'] = self.process_id
|
||||
await workspace.box_service.start_managed_process(workspace.session_id, payload)
|
||||
except Exception:
|
||||
self.owner.error_phase = MCPSessionErrorPhase.PROCESS_START
|
||||
raise
|
||||
else:
|
||||
self.ap.logger.info(
|
||||
f'MCP server {self.server_name}: reusing live managed process '
|
||||
f'process_id={self.process_id} (transport reconnect)'
|
||||
)
|
||||
|
||||
websocket_url = workspace.get_managed_process_websocket_url(self.process_id)
|
||||
|
||||
# Attach the WS transport + MCP session ONCE, on the owner's exit stack,
|
||||
# in the same task as the serve loop that follows. websocket_client and
|
||||
# ClientSession use anyio task groups whose cancel scope is bound to the
|
||||
# frame/stack that entered them, so they must live on the owner exit
|
||||
# stack (not a deferred/transferred one) or the streams close the moment
|
||||
# initialize() returns and the next request fails with "Connection
|
||||
# closed".
|
||||
#
|
||||
# A slow (`npx -y <pkg>`) cold start makes this single attempt fail
|
||||
# while the process is still alive — the package is still installing and
|
||||
# cannot answer the handshake. We surface that to the outer retry loop
|
||||
# as a _ColdStartRetry: it must NOT stop the process (it is healthy and
|
||||
# will be reused) and must NOT consume the fatal retry budget. The next
|
||||
# attempt re-attaches to the same live process; once it has finished
|
||||
# cold start the handshake succeeds and stays healthy.
|
||||
try:
|
||||
process_workspace = (
|
||||
self._build_workspace(host_path=host_path, workdir=process_cwd, mount_path=process_cwd)
|
||||
if host_path
|
||||
else workspace
|
||||
)
|
||||
payload = process_workspace.build_process_payload(
|
||||
self.server_config['command'],
|
||||
self.server_config.get('args', []),
|
||||
env=self.server_config.get('env', {}),
|
||||
cwd=process_cwd,
|
||||
)
|
||||
if install_cmd:
|
||||
payload = self._wrap_process_payload_with_python_env(payload, process_cwd)
|
||||
payload['process_id'] = self.process_id
|
||||
await workspace.box_service.start_managed_process(workspace.session_id, payload)
|
||||
except Exception:
|
||||
self.owner.error_phase = MCPSessionErrorPhase.PROCESS_START
|
||||
raise
|
||||
|
||||
try:
|
||||
websocket_url = workspace.get_managed_process_websocket_url(self.process_id)
|
||||
transport = await self.owner.exit_stack.enter_async_context(websocket_client(websocket_url))
|
||||
read_stream, write_stream = transport
|
||||
self.owner.session = await self.owner.exit_stack.enter_async_context(
|
||||
@@ -285,19 +202,12 @@ class BoxStdioSessionRuntime:
|
||||
)
|
||||
except Exception:
|
||||
self.owner.error_phase = MCPSessionErrorPhase.RELAY_CONNECT
|
||||
if not await self._managed_process_has_exited():
|
||||
# Process is alive but not yet serving (cold start) — reconnect.
|
||||
raise _ColdStartRetry(f'{self.server_name}: transport not ready during cold start')
|
||||
raise
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(self.owner.session.initialize(), timeout=_HANDSHAKE_ATTEMPT_TIMEOUT_SEC)
|
||||
except Exception as exc:
|
||||
await self.owner.session.initialize()
|
||||
except Exception:
|
||||
self.owner.error_phase = MCPSessionErrorPhase.MCP_INIT
|
||||
if not await self._managed_process_has_exited():
|
||||
raise _ColdStartRetry(
|
||||
f'{self.server_name}: handshake not ready during cold start ({type(exc).__name__})'
|
||||
)
|
||||
raise
|
||||
|
||||
async def monitor_process_health(self) -> None:
|
||||
@@ -324,74 +234,8 @@ class BoxStdioSessionRuntime:
|
||||
)
|
||||
if consecutive_errors >= self.owner._MONITOR_MAX_CONSECUTIVE_ERRORS:
|
||||
return
|
||||
|
||||
# Capture stderr logs from the managed process
|
||||
if isinstance(info, dict):
|
||||
stderr_text = info.get('stderr', '') or info.get('stderr_preview', '')
|
||||
else:
|
||||
stderr_text = getattr(info, 'stderr', '') or getattr(info, 'stderr_preview', '')
|
||||
|
||||
if stderr_text and stderr_text != self.owner._last_stderr_text:
|
||||
# Find new lines not in the previous snapshot
|
||||
old_lines = set(self.owner._last_stderr_text.splitlines()) if self.owner._last_stderr_text else set()
|
||||
new_lines = [l for l in stderr_text.splitlines() if l and l not in old_lines]
|
||||
self.owner._last_stderr_text = stderr_text
|
||||
|
||||
import time as _time
|
||||
|
||||
for line in new_lines:
|
||||
level = (
|
||||
'error'
|
||||
if any(k in line.upper() for k in ('ERROR', 'CRITICAL'))
|
||||
else 'warning'
|
||||
if 'WARNING' in line.upper()
|
||||
else 'debug'
|
||||
if 'DEBUG' in line.upper()
|
||||
else 'info'
|
||||
)
|
||||
self.owner._log_buffer.append({'ts': _time.time(), 'level': level, 'text': line})
|
||||
|
||||
await asyncio.sleep(self.owner._MONITOR_POLL_INTERVAL)
|
||||
|
||||
async def _managed_process_is_running(self) -> bool:
|
||||
"""Return True if this server's managed process exists and is running.
|
||||
|
||||
Used to decide whether initialize() must (re)start the process or can
|
||||
simply re-attach the WebSocket transport to a process the Box runtime
|
||||
kept alive across a transient transport drop.
|
||||
"""
|
||||
from langbot_plugin.box.models import BoxManagedProcessStatus
|
||||
|
||||
workspace = self._build_workspace()
|
||||
try:
|
||||
info = await workspace.get_managed_process(self.process_id)
|
||||
except Exception:
|
||||
return False
|
||||
status = info.get('status', '') if isinstance(info, dict) else getattr(info, 'status', '')
|
||||
return status in (BoxManagedProcessStatus.RUNNING.value, BoxManagedProcessStatus.RUNNING)
|
||||
|
||||
async def _managed_process_has_exited(self) -> bool:
|
||||
"""Return True only if the process is DEFINITIVELY gone (reports EXITED).
|
||||
|
||||
Distinct from ``not _managed_process_is_running()``: a process that has
|
||||
just been spawned may not yet report RUNNING, and a transient query
|
||||
error is not proof of exit. During the cold-start handshake retry we
|
||||
must NOT treat 'not yet running' or 'query failed' as a terminal
|
||||
failure, or we bail out to the outer rebuild path and churn the
|
||||
process (relay then rejects the early re-attach with HTTP 400). Only a
|
||||
successful query that reports EXITED stops the retry loop.
|
||||
"""
|
||||
from langbot_plugin.box.models import BoxManagedProcessStatus
|
||||
|
||||
workspace = self._build_workspace()
|
||||
try:
|
||||
info = await workspace.get_managed_process(self.process_id)
|
||||
except Exception:
|
||||
# Unknown — treat as 'still coming up', not exited.
|
||||
return False
|
||||
status = info.get('status', '') if isinstance(info, dict) else getattr(info, 'status', '')
|
||||
return status in (BoxManagedProcessStatus.EXITED.value, BoxManagedProcessStatus.EXITED)
|
||||
|
||||
async def _stage_host_path_to_shared_workspace(self, host_path: str) -> str:
|
||||
source_path = normalize_host_path(host_path)
|
||||
if not source_path:
|
||||
@@ -498,20 +342,16 @@ class BoxStdioSessionRuntime:
|
||||
|
||||
workspace = self._build_workspace(host_path=None)
|
||||
|
||||
# Transient config-page tests now share the same 'mcp-shared' Box
|
||||
# session as live servers, so we must NOT tear the session down here —
|
||||
# that would kill every other MCP server in the container. A test is
|
||||
# isolated at the process level: it ran under its own process_id, so we
|
||||
# stop only that process, exactly like a live server does below. The
|
||||
# shared session and all other servers' live processes are untouched.
|
||||
# (Staged per-test workspace files are still cleaned up.)
|
||||
# Transient test sessions own their isolated Box session, so tear the
|
||||
# whole session down rather than leaking it. This cannot affect live
|
||||
# servers because they live in the separate shared session.
|
||||
if getattr(self.owner, 'is_transient', False):
|
||||
try:
|
||||
await workspace.stop_managed_process(self.process_id)
|
||||
await workspace.cleanup()
|
||||
except Exception as exc:
|
||||
self.ap.logger.warning(
|
||||
f'MCP server {self.server_name}: failed to stop transient test process '
|
||||
f'process_id={self.process_id}: {type(exc).__name__}: {exc}'
|
||||
f'MCP server {self.server_name}: failed to delete transient test session '
|
||||
f'{self.owner._build_box_session_id()}: {type(exc).__name__}: {exc}'
|
||||
)
|
||||
await self._cleanup_staged_workspace()
|
||||
return
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import typing
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
|
||||
@@ -143,130 +142,21 @@ class ToolManager:
|
||||
|
||||
return tools
|
||||
|
||||
def _get_query_session_id(self, query: pipeline_query.Query) -> str | None:
|
||||
launcher_type = getattr(query, 'launcher_type', None)
|
||||
launcher_id = getattr(query, 'launcher_id', None)
|
||||
if launcher_type is None or launcher_id is None:
|
||||
return None
|
||||
|
||||
launcher_type_value = launcher_type.value if hasattr(launcher_type, 'value') else launcher_type
|
||||
return f'{launcher_type_value}_{launcher_id}'
|
||||
|
||||
async def _record_tool_call(
|
||||
self,
|
||||
*,
|
||||
name: str,
|
||||
source: str,
|
||||
parameters: dict,
|
||||
query: pipeline_query.Query,
|
||||
duration_ms: int,
|
||||
status: str,
|
||||
result: typing.Any = None,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
monitoring_service = getattr(self.ap, 'monitoring_service', None)
|
||||
if not monitoring_service:
|
||||
return
|
||||
|
||||
variables = getattr(query, 'variables', {}) or {}
|
||||
message_id = variables.get('_monitoring_message_id') if isinstance(variables, dict) else None
|
||||
bot_name = variables.get('_monitoring_bot_name') if isinstance(variables, dict) else None
|
||||
pipeline_name = variables.get('_monitoring_pipeline_name') if isinstance(variables, dict) else None
|
||||
|
||||
try:
|
||||
await monitoring_service.record_tool_call(
|
||||
tool_name=name,
|
||||
tool_source=source,
|
||||
duration=duration_ms,
|
||||
status=status,
|
||||
bot_id=getattr(query, 'bot_uuid', None),
|
||||
bot_name=bot_name,
|
||||
pipeline_name=pipeline_name,
|
||||
session_id=self._get_query_session_id(query),
|
||||
message_id=message_id,
|
||||
arguments=parameters,
|
||||
result=result,
|
||||
error_message=error_message,
|
||||
)
|
||||
except Exception as e:
|
||||
self.ap.logger.warning(f'Failed to record tool call: {e}')
|
||||
|
||||
async def _invoke_tool_with_monitoring(
|
||||
self,
|
||||
*,
|
||||
source: str,
|
||||
name: str,
|
||||
parameters: dict,
|
||||
query: pipeline_query.Query,
|
||||
invoke: typing.Callable[[], typing.Awaitable[typing.Any]],
|
||||
) -> typing.Any:
|
||||
start_time = time.perf_counter()
|
||||
try:
|
||||
result = await invoke()
|
||||
except Exception as e:
|
||||
duration_ms = int((time.perf_counter() - start_time) * 1000)
|
||||
await self._record_tool_call(
|
||||
name=name,
|
||||
source=source,
|
||||
parameters=parameters,
|
||||
query=query,
|
||||
duration_ms=duration_ms,
|
||||
status='error',
|
||||
error_message=str(e),
|
||||
)
|
||||
raise
|
||||
|
||||
duration_ms = int((time.perf_counter() - start_time) * 1000)
|
||||
await self._record_tool_call(
|
||||
name=name,
|
||||
source=source,
|
||||
parameters=parameters,
|
||||
query=query,
|
||||
duration_ms=duration_ms,
|
||||
status='success',
|
||||
result=result,
|
||||
)
|
||||
return result
|
||||
|
||||
async def execute_func_call(self, name: str, parameters: dict, query: pipeline_query.Query) -> typing.Any:
|
||||
from langbot.pkg.telemetry import features as telemetry_features
|
||||
|
||||
if await self.native_tool_loader.has_tool(name):
|
||||
telemetry_features.increment(query, 'tool_calls', 'native')
|
||||
return await self._invoke_tool_with_monitoring(
|
||||
source='native',
|
||||
name=name,
|
||||
parameters=parameters,
|
||||
query=query,
|
||||
invoke=lambda: self.native_tool_loader.invoke_tool(name, parameters, query),
|
||||
)
|
||||
return await self.native_tool_loader.invoke_tool(name, parameters, query)
|
||||
if await self.plugin_tool_loader.has_tool(name):
|
||||
telemetry_features.increment(query, 'tool_calls', 'plugin')
|
||||
return await self._invoke_tool_with_monitoring(
|
||||
source='plugin',
|
||||
name=name,
|
||||
parameters=parameters,
|
||||
query=query,
|
||||
invoke=lambda: self.plugin_tool_loader.invoke_tool(name, parameters, query),
|
||||
)
|
||||
return await self.plugin_tool_loader.invoke_tool(name, parameters, query)
|
||||
if await self.mcp_tool_loader.has_tool(name):
|
||||
telemetry_features.increment(query, 'tool_calls', 'mcp')
|
||||
return await self._invoke_tool_with_monitoring(
|
||||
source='mcp',
|
||||
name=name,
|
||||
parameters=parameters,
|
||||
query=query,
|
||||
invoke=lambda: self.mcp_tool_loader.invoke_tool(name, parameters, query),
|
||||
)
|
||||
return await self.mcp_tool_loader.invoke_tool(name, parameters, query)
|
||||
if await self.skill_tool_loader.has_tool(name):
|
||||
telemetry_features.increment(query, 'tool_calls', 'skill')
|
||||
return await self._invoke_tool_with_monitoring(
|
||||
source='skill',
|
||||
name=name,
|
||||
parameters=parameters,
|
||||
query=query,
|
||||
invoke=lambda: self.skill_tool_loader.invoke_tool(name, parameters, query),
|
||||
)
|
||||
return await self.skill_tool_loader.invoke_tool(name, parameters, query)
|
||||
raise ToolNotFoundError(name)
|
||||
|
||||
async def shutdown(self):
|
||||
|
||||
@@ -33,12 +33,6 @@ class VectorDBManager:
|
||||
self.vector_db = SeekDBVectorDatabase(self.ap)
|
||||
self.ap.logger.info('Initialized SeekDB vector database backend.')
|
||||
|
||||
elif vdb_type == 'valkey_search':
|
||||
from .vdbs.valkey_search import ValkeySearchVectorDatabase
|
||||
|
||||
self.vector_db = ValkeySearchVectorDatabase(self.ap)
|
||||
self.ap.logger.info('Initialized Valkey Search vector database backend.')
|
||||
|
||||
elif vdb_type == 'milvus':
|
||||
from .vdbs.milvus import MilvusVectorDatabase
|
||||
|
||||
|
||||
@@ -1,828 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import struct
|
||||
from typing import Any
|
||||
|
||||
from langbot.pkg.core import app
|
||||
from langbot.pkg.vector.vdb import VectorDatabase, SearchType
|
||||
from langbot.pkg.vector.filter_utils import normalize_filter, strip_unsupported_fields
|
||||
|
||||
try:
|
||||
from glide import (
|
||||
Batch,
|
||||
GlideClient,
|
||||
GlideClientConfiguration,
|
||||
NodeAddress,
|
||||
RequestError,
|
||||
ServerCredentials,
|
||||
ft,
|
||||
VectorField,
|
||||
VectorFieldAttributesHnsw,
|
||||
VectorFieldAttributesFlat,
|
||||
VectorAlgorithm,
|
||||
VectorType,
|
||||
DistanceMetricType,
|
||||
TagField,
|
||||
TextField,
|
||||
FtCreateOptions,
|
||||
DataType,
|
||||
FtSearchOptions,
|
||||
FtSearchLimit,
|
||||
ReturnField,
|
||||
)
|
||||
|
||||
VALKEY_SEARCH_AVAILABLE = True
|
||||
except ImportError:
|
||||
VALKEY_SEARCH_AVAILABLE = False
|
||||
|
||||
# Default per-request timeout (ms) for the glide client. The glide library
|
||||
# default is 250ms, which is too low for vector KNN (``FT.SEARCH ... =>[KNN]``)
|
||||
# under moderate load or with large indexes and yields spurious TimeoutErrors.
|
||||
# Overridable via the ``vdb.valkey_search.request_timeout`` config option.
|
||||
_DEFAULT_REQUEST_TIMEOUT_MS = 5000
|
||||
|
||||
# Safety cap on the number of SCAN rounds when purging a collection's keys, so
|
||||
# a cursor-handling bug or pathological keyspace can never spin forever.
|
||||
_MAX_SCAN_ROUNDS = 100000
|
||||
|
||||
|
||||
# Mandatory client name for production observability (CLIENT LIST / dashboards).
|
||||
VALKEY_CLIENT_NAME = 'langbot_vector_client'
|
||||
|
||||
# Fixed, indexed metadata schema. LangBot's RAG layer stores ``file_id`` on
|
||||
# every chunk; it is the only metadata field we promote to a first-class
|
||||
# (filterable) index field. All other metadata is preserved verbatim inside
|
||||
# the ``metadata_json`` field so it survives a round-trip, but is NOT
|
||||
# filterable (the established Milvus / pgvector pragmatism).
|
||||
_INDEXED_TAG_FIELDS = {'file_id'}
|
||||
_SUPPORTED_FILTER_FIELDS = set(_INDEXED_TAG_FIELDS)
|
||||
|
||||
# Hash field names used for stored documents.
|
||||
_FIELD_VECTOR = 'vector'
|
||||
_FIELD_DOCUMENT = 'document'
|
||||
_FIELD_FILE_ID = 'file_id'
|
||||
_FIELD_METADATA = 'metadata_json'
|
||||
_VEC_SCORE_ALIAS = '__vec_score'
|
||||
|
||||
# Valkey Search has no bare "match everything" token for non-vector queries
|
||||
# (a standalone ``*`` is a syntax error). A negated match on a sentinel tag
|
||||
# value that can never exist matches every key, which is the canonical
|
||||
# match-all idiom for FT.SEARCH.
|
||||
_MATCH_ALL = '-@file_id:{__langbot_match_all_sentinel__}'
|
||||
|
||||
# Page size used when enumerating matching keys for deletion. Deletes
|
||||
# paginate through the full result set in batches of this size so that
|
||||
# files/filters matching more than one page of chunks are fully removed
|
||||
# (no silent truncation / orphaned vectors).
|
||||
_DELETE_SCAN_BATCH = 10000
|
||||
|
||||
# Characters Valkey Search's TAG query parser cannot handle even when
|
||||
# backslash-escaped (the brace delimiters and the wildcard). file_id TAG
|
||||
# values are percent-encoded over this set (plus '%' itself, so the encoding
|
||||
# is reversible/unambiguous) before being stored or queried, so an arbitrary
|
||||
# file_id round-trips instead of producing an unparseable query. For normal
|
||||
# UUID/hash file_ids none of these characters occur, so the encoding is a
|
||||
# no-op and the stored value is unchanged. The original file_id is always
|
||||
# preserved verbatim inside ``metadata_json``.
|
||||
_FT_UNSAFE_TAG_CHARS = frozenset('{}*%')
|
||||
|
||||
|
||||
class ValkeySearchVectorDatabase(VectorDatabase):
|
||||
"""Valkey Search (valkey-bundle) vector database adapter for LangBot.
|
||||
|
||||
Backed by the Valkey Search module shipped in ``valkey/valkey-bundle``,
|
||||
accessed through the official ``valkey-glide`` client's native ``ft``
|
||||
(search) command namespace. Documents are stored as Valkey HASH keys
|
||||
under a per-collection prefix and indexed by one ``FT.CREATE`` index per
|
||||
collection.
|
||||
|
||||
Supported search types: ``VECTOR``, ``FULL_TEXT`` and ``HYBRID``.
|
||||
|
||||
Hybrid search semantics (IMPORTANT)
|
||||
-----------------------------------
|
||||
Valkey Search hybrid queries follow a *filter-then-KNN* model: the text /
|
||||
metadata filter pre-selects candidate keys and the KNN stage ranks them by
|
||||
vector distance. This backend does **NOT** implement application-side
|
||||
weighted score fusion. The ``vector_weight`` argument is therefore
|
||||
accepted for interface compatibility but is **not honored** — passing
|
||||
different weights does not change result ordering. A one-time warning is
|
||||
emitted the first time a non-default weight is supplied. App-side score
|
||||
fusion can be layered on later if weighted hybrid ranking is required.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def supported_search_types(cls) -> list[SearchType]:
|
||||
return [SearchType.VECTOR, SearchType.FULL_TEXT, SearchType.HYBRID]
|
||||
|
||||
def __init__(self, ap: app.Application):
|
||||
if not VALKEY_SEARCH_AVAILABLE:
|
||||
raise ImportError(
|
||||
"valkey-glide is not installed. Install it with: pip install 'valkey-glide>=2.4.1,<3.0.0'"
|
||||
)
|
||||
|
||||
self.ap = ap
|
||||
config = self.ap.instance_config.data['vdb']['valkey_search']
|
||||
|
||||
self._host = config.get('host', 'localhost')
|
||||
self._port = int(config.get('port', 6379))
|
||||
self._db = int(config.get('db', 0))
|
||||
# Auth / TLS are optional (toB / SaaS). Never logged.
|
||||
self._password = config.get('password', '') or None
|
||||
self._username = config.get('username', '') or None
|
||||
self._tls = bool(config.get('tls', False))
|
||||
self._request_timeout = int(config.get('request_timeout', _DEFAULT_REQUEST_TIMEOUT_MS))
|
||||
|
||||
algorithm = str(config.get('index_algorithm', 'HNSW')).upper()
|
||||
self._algorithm = VectorAlgorithm.FLAT if algorithm == 'FLAT' else VectorAlgorithm.HNSW
|
||||
|
||||
metric = str(config.get('distance_metric', 'COSINE')).upper()
|
||||
self._distance_metric = {
|
||||
'COSINE': DistanceMetricType.COSINE,
|
||||
'L2': DistanceMetricType.L2,
|
||||
'IP': DistanceMetricType.IP,
|
||||
}.get(metric, DistanceMetricType.COSINE)
|
||||
|
||||
# Lazily-created client (created on first use so a down Valkey does not
|
||||
# block LangBot boot).
|
||||
self._client: GlideClient | None = None
|
||||
# Serializes lazy client creation so concurrent first-use callers do not
|
||||
# each construct (and leak) a separate GlideClient.
|
||||
self._client_lock = asyncio.Lock()
|
||||
# Index names we have already ensured this process lifetime.
|
||||
self._ensured_indexes: set[str] = set()
|
||||
# Whether we have already warned about the non-honored vector_weight.
|
||||
self._vector_weight_warned = False
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Client lifecycle
|
||||
# ------------------------------------------------------------------ #
|
||||
async def _ensure_client(self) -> GlideClient:
|
||||
"""Create the glide client on first use (lazy, non-blocking boot)."""
|
||||
if self._client is not None:
|
||||
return self._client
|
||||
# Double-checked locking: serialize creation so two concurrent
|
||||
# first-use callers don't both build a client and leak one.
|
||||
async with self._client_lock:
|
||||
if self._client is not None:
|
||||
return self._client
|
||||
|
||||
credentials = None
|
||||
if self._password is not None:
|
||||
# username is optional alongside a password (ACL "user" vs default user).
|
||||
credentials = ServerCredentials(password=self._password, username=self._username)
|
||||
elif self._username is not None:
|
||||
# A username without a password is not a valid credential pair, and silently
|
||||
# connecting unauthenticated to a potentially shared Valkey instance is a
|
||||
# security footgun (e.g. an env var that failed to resolve). Fail closed.
|
||||
raise ValueError(
|
||||
'Valkey Search: a username was configured without a password. '
|
||||
'Set both username and password to use ACL authentication, or remove both.'
|
||||
)
|
||||
|
||||
conf = GlideClientConfiguration(
|
||||
addresses=[NodeAddress(self._host, self._port)],
|
||||
client_name=VALKEY_CLIENT_NAME,
|
||||
database_id=self._db,
|
||||
use_tls=self._tls,
|
||||
lazy_connect=True,
|
||||
credentials=credentials,
|
||||
request_timeout=self._request_timeout,
|
||||
)
|
||||
self._client = await GlideClient.create(conf)
|
||||
self.ap.logger.info(
|
||||
f'Initialized Valkey Search client to {self._host}:{self._port} (db={self._db}, tls={self._tls})'
|
||||
)
|
||||
return self._client
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the glide client and reset state.
|
||||
|
||||
Safe to call when no client was created. After ``close`` the next
|
||||
operation transparently re-creates the client (``_ensure_client``
|
||||
guards on ``self._client is None``).
|
||||
"""
|
||||
if self._client is not None:
|
||||
try:
|
||||
await self._client.close()
|
||||
except Exception:
|
||||
self.ap.logger.warning('Valkey Search: error while closing client (ignored)')
|
||||
finally:
|
||||
self._client = None
|
||||
self._ensured_indexes.clear()
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Naming helpers
|
||||
# ------------------------------------------------------------------ #
|
||||
@staticmethod
|
||||
def _index_name(collection: str) -> str:
|
||||
return f'idx:{collection}'
|
||||
|
||||
@staticmethod
|
||||
def _key_prefix(collection: str) -> str:
|
||||
return f'kb:{collection}:'
|
||||
|
||||
@staticmethod
|
||||
def _pack_vector(vec: list[float]) -> bytes:
|
||||
"""Pack a float vector into little-endian float32 bytes.
|
||||
|
||||
Valkey Search stores and queries vectors as FLOAT32 little-endian
|
||||
blobs (per the search query-language spec).
|
||||
"""
|
||||
return struct.pack(f'<{len(vec)}f', *[float(x) for x in vec])
|
||||
|
||||
@staticmethod
|
||||
def _escape_tag(value: str) -> str:
|
||||
"""Escape characters that are special inside a TAG ``{...}`` clause.
|
||||
|
||||
The backslash is escaped first so it cannot consume a following
|
||||
escape. This neutralises injection-style values (quotes, parens,
|
||||
``|``, ``@``, ``:``, spaces, dashes) so a crafted ``file_id`` cannot
|
||||
break out of the clause.
|
||||
|
||||
Note: Valkey Search's TAG query parser cannot handle a literal brace
|
||||
(``{`` / ``}``) or ``*`` even when backslash-escaped. Callers that pass
|
||||
a ``file_id`` route it through ``_encode_and_escape_tag`` /
|
||||
``_encode_file_id`` first, which percent-encodes exactly those
|
||||
characters, so an arbitrary ``file_id`` round-trips safely. This raw
|
||||
escaper is still correct for all other special characters.
|
||||
"""
|
||||
out = []
|
||||
for ch in str(value):
|
||||
if ch in '\\,.<>{}[]"\':;!@#$%^&*()-+=~| ':
|
||||
out.append('\\')
|
||||
out.append(ch)
|
||||
return ''.join(out)
|
||||
|
||||
@staticmethod
|
||||
def _encode_file_id(value: str) -> str:
|
||||
"""Make a ``file_id`` safe to use as an FT TAG token AND query value.
|
||||
|
||||
Percent-encodes the characters Valkey Search's TAG parser cannot handle
|
||||
even when backslash-escaped (``{``, ``}``, ``*``) plus ``%`` itself for
|
||||
reversibility. Applied identically at write time (the stored TAG field)
|
||||
and query time (filters / ``delete_by_file_id``) so any value matches
|
||||
itself. For normal UUID/hash ids none of these characters occur, so
|
||||
this is a no-op. The original value is always kept verbatim in
|
||||
``metadata_json``; this encoded form is only ever used for the indexed
|
||||
TAG.
|
||||
"""
|
||||
out = []
|
||||
for ch in str(value):
|
||||
if ch in _FT_UNSAFE_TAG_CHARS:
|
||||
out.append('%{:02X}'.format(ord(ch)))
|
||||
else:
|
||||
out.append(ch)
|
||||
return ''.join(out)
|
||||
|
||||
def _encode_and_escape_tag(self, value: str) -> str:
|
||||
"""Encode an FT-unsafe ``file_id`` then escape TAG special chars."""
|
||||
return self._escape_tag(self._encode_file_id(value))
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Filter mapping (canonical triples -> FT query fragment)
|
||||
# ------------------------------------------------------------------ #
|
||||
def _triples_to_ft(self, filter: dict[str, Any] | None) -> str:
|
||||
"""Translate a canonical filter dict into an FT filter expression.
|
||||
|
||||
Only indexed fields (``file_id``) are filterable; unsupported fields
|
||||
are dropped with a warning (matching the Milvus / pgvector pattern).
|
||||
Returns an empty string when there is no usable filter.
|
||||
"""
|
||||
triples = normalize_filter(filter)
|
||||
if not triples:
|
||||
return ''
|
||||
triples = strip_unsupported_fields(triples, _SUPPORTED_FILTER_FIELDS)
|
||||
|
||||
fragments: list[str] = []
|
||||
for field, op, value in triples:
|
||||
# All currently-indexed fields are TAG fields; file_id values are
|
||||
# encoded (FT-unsafe chars) then escaped so any value round-trips.
|
||||
if op == '$eq':
|
||||
fragments.append(f'@{field}:{{{self._encode_and_escape_tag(value)}}}')
|
||||
elif op == '$ne':
|
||||
fragments.append(f'-@{field}:{{{self._encode_and_escape_tag(value)}}}')
|
||||
elif op == '$in':
|
||||
joined = '|'.join(self._encode_and_escape_tag(v) for v in value)
|
||||
fragments.append(f'@{field}:{{{joined}}}')
|
||||
elif op == '$nin':
|
||||
joined = '|'.join(self._encode_and_escape_tag(v) for v in value)
|
||||
fragments.append(f'-@{field}:{{{joined}}}')
|
||||
elif op == '$gt':
|
||||
fragments.append(f'@{field}:[({float(value)} +inf]')
|
||||
elif op == '$gte':
|
||||
fragments.append(f'@{field}:[{float(value)} +inf]')
|
||||
elif op == '$lt':
|
||||
fragments.append(f'@{field}:[-inf ({float(value)}]')
|
||||
elif op == '$lte':
|
||||
fragments.append(f'@{field}:[-inf {float(value)}]')
|
||||
else:
|
||||
# normalize_filter() already rejects unknown operators, so this
|
||||
# only triggers if SUPPORTED_OPS grows without this chain being
|
||||
# updated. Fail closed (rather than silently dropping the
|
||||
# condition, which would widen delete_by_filter's match set).
|
||||
raise ValueError(f'Valkey Search: unhandled filter operator {op!r} on field {field!r}')
|
||||
|
||||
return ' '.join(fragments)
|
||||
|
||||
@staticmethod
|
||||
def _build_text_clause(text: str) -> str:
|
||||
"""Build a field-scoped full-text clause for the ``document`` field.
|
||||
|
||||
Each whitespace-delimited word becomes a ``@document:<term>`` term and
|
||||
the terms are AND-ed (space separated). FT special characters in each
|
||||
term are escaped. Returns an empty string when *text* has no words.
|
||||
"""
|
||||
words = [w for w in str(text).split() if w]
|
||||
if not words:
|
||||
return ''
|
||||
terms = [f'@{_FIELD_DOCUMENT}:{ValkeySearchVectorDatabase._escape_text(w)}' for w in words]
|
||||
return ' '.join(terms)
|
||||
|
||||
@staticmethod
|
||||
def _escape_text(text: str) -> str:
|
||||
"""Escape FT full-text special characters in a single term."""
|
||||
out = []
|
||||
for ch in str(text):
|
||||
if ch in '@!{}[]()|-"~*:\\':
|
||||
out.append('\\')
|
||||
out.append(ch)
|
||||
return ''.join(out)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Index management
|
||||
# ------------------------------------------------------------------ #
|
||||
async def _ensure_index(self, client: GlideClient, collection: str, dim: int) -> None:
|
||||
index = self._index_name(collection)
|
||||
if index in self._ensured_indexes:
|
||||
return
|
||||
|
||||
# ft.info is O(1) and raises RequestError when the index is absent —
|
||||
# cheaper than ft.list (O(n) over all indexes) and it closes the
|
||||
# check-then-create TOCTOU window.
|
||||
try:
|
||||
await ft.info(client, index)
|
||||
self._ensured_indexes.add(index)
|
||||
return
|
||||
except RequestError:
|
||||
pass
|
||||
|
||||
if self._algorithm == VectorAlgorithm.FLAT:
|
||||
vector_attrs = VectorFieldAttributesFlat(
|
||||
dimensions=dim,
|
||||
distance_metric=self._distance_metric,
|
||||
type=VectorType.FLOAT32,
|
||||
)
|
||||
else:
|
||||
vector_attrs = VectorFieldAttributesHnsw(
|
||||
dimensions=dim,
|
||||
distance_metric=self._distance_metric,
|
||||
type=VectorType.FLOAT32,
|
||||
)
|
||||
|
||||
schema = [
|
||||
VectorField(name=_FIELD_VECTOR, algorithm=self._algorithm, attributes=vector_attrs),
|
||||
TagField(name=_FIELD_FILE_ID),
|
||||
TextField(name=_FIELD_DOCUMENT),
|
||||
]
|
||||
options = FtCreateOptions(data_type=DataType.HASH, prefixes=[self._key_prefix(collection)])
|
||||
await ft.create(client, index, schema, options)
|
||||
self._ensured_indexes.add(index)
|
||||
self.ap.logger.info(
|
||||
f"Valkey Search index '{index}' created (dim={dim}, algo={self._algorithm.value}, "
|
||||
f'metric={self._distance_metric.value})'
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _decode(value: Any) -> str:
|
||||
if isinstance(value, (bytes, bytearray, memoryview)):
|
||||
return bytes(value).decode('utf-8', errors='replace')
|
||||
return str(value)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# VectorDatabase ABC implementation
|
||||
# ------------------------------------------------------------------ #
|
||||
async def get_or_create_collection(self, collection: str):
|
||||
"""Ensure a client exists.
|
||||
|
||||
The index itself requires the vector dimension, which is only known at
|
||||
first ``add_embeddings`` (same constraint as Qdrant / SeekDB), so this
|
||||
is a best-effort no-op when the index does not yet exist.
|
||||
"""
|
||||
await self._ensure_client()
|
||||
|
||||
async def add_embeddings(
|
||||
self,
|
||||
collection: str,
|
||||
ids: list[str],
|
||||
embeddings_list: list[list[float]],
|
||||
metadatas: list[dict[str, Any]],
|
||||
documents: list[str] | None = None,
|
||||
) -> None:
|
||||
if not embeddings_list:
|
||||
return
|
||||
|
||||
client = await self._ensure_client()
|
||||
dim = len(embeddings_list[0])
|
||||
# The index schema is fixed to the first embedding's dimension. A later
|
||||
# embedding of a different length would be packed into a wrong-sized
|
||||
# blob that Valkey stores silently but that yields garbage KNN
|
||||
# distances, so reject mixed dimensions up-front.
|
||||
if any(len(e) != dim for e in embeddings_list[1:]):
|
||||
raise ValueError(f'All embeddings must have dimension {dim}; got mixed lengths')
|
||||
await self._ensure_index(client, collection, dim)
|
||||
|
||||
prefix = self._key_prefix(collection)
|
||||
|
||||
batch = Batch(is_atomic=False)
|
||||
for i, _id in enumerate(ids):
|
||||
key = prefix + str(_id)
|
||||
metadata = metadatas[i] if i < len(metadatas) else {}
|
||||
mapping: dict[str, Any] = {
|
||||
_FIELD_VECTOR: self._pack_vector(embeddings_list[i]),
|
||||
_FIELD_METADATA: json.dumps(metadata, ensure_ascii=False),
|
||||
}
|
||||
file_id = metadata.get('file_id')
|
||||
if file_id is not None:
|
||||
mapping[_FIELD_FILE_ID] = self._encode_file_id(str(file_id))
|
||||
if documents is not None and i < len(documents) and documents[i] is not None:
|
||||
mapping[_FIELD_DOCUMENT] = documents[i]
|
||||
|
||||
batch.hset(key, mapping)
|
||||
|
||||
# Pipeline all HSETs into a single round-trip (non-atomic) instead of
|
||||
# one await per embedding, which is N sequential round-trips for N
|
||||
# chunks.
|
||||
await client.exec(batch, raise_on_error=True)
|
||||
|
||||
self.ap.logger.info(f"Added {len(ids)} embeddings to Valkey Search collection '{collection}'")
|
||||
|
||||
async def search(
|
||||
self,
|
||||
collection: str,
|
||||
query_embedding: list[float],
|
||||
k: int = 5,
|
||||
search_type: str = 'vector',
|
||||
query_text: str = '',
|
||||
filter: dict[str, Any] | None = None,
|
||||
vector_weight: float | None = None,
|
||||
) -> dict[str, Any]:
|
||||
client = await self._ensure_client()
|
||||
index = self._index_name(collection)
|
||||
|
||||
if not await self._index_exists(client, index):
|
||||
return {'ids': [[]], 'metadatas': [[]], 'distances': [[]]}
|
||||
|
||||
# vector_weight is accepted for interface parity but NOT honored by this
|
||||
# backend (filter-then-KNN, no weighted fusion). Warn once.
|
||||
if vector_weight is not None and not self._vector_weight_warned:
|
||||
self.ap.logger.warning(
|
||||
'Valkey Search backend does not honor vector_weight: hybrid search uses '
|
||||
'filter-then-KNN without weighted score fusion. The vector_weight value '
|
||||
'is ignored. See docs/VALKEY_SEARCH_INTEGRATION.md.'
|
||||
)
|
||||
self._vector_weight_warned = True
|
||||
|
||||
filter_expr = self._triples_to_ft(filter)
|
||||
|
||||
if search_type == SearchType.FULL_TEXT:
|
||||
if not query_text:
|
||||
return {'ids': [[]], 'metadatas': [[]], 'distances': [[]]}
|
||||
text_clause = self._build_text_clause(query_text)
|
||||
if not text_clause:
|
||||
return {'ids': [[]], 'metadatas': [[]], 'distances': [[]]}
|
||||
query = f'{filter_expr} {text_clause}'.strip() if filter_expr else text_clause
|
||||
return await self._run_text_search(client, index, query, k)
|
||||
|
||||
if search_type == SearchType.HYBRID:
|
||||
# Filter / text pre-selects candidates; KNN ranks. No fusion.
|
||||
pre = filter_expr
|
||||
if query_text:
|
||||
text_clause = self._build_text_clause(query_text)
|
||||
if text_clause:
|
||||
pre = f'{pre} {text_clause}'.strip() if pre else text_clause
|
||||
pre = pre or '*'
|
||||
query = f'{self._wrap_pre(pre)}=>[KNN {k} @{_FIELD_VECTOR} $BLOB AS {_VEC_SCORE_ALIAS}]'
|
||||
return await self._run_knn_search(client, index, query, query_embedding, k)
|
||||
|
||||
# Default: pure VECTOR search.
|
||||
pre = filter_expr or '*'
|
||||
query = f'{self._wrap_pre(pre)}=>[KNN {k} @{_FIELD_VECTOR} $BLOB AS {_VEC_SCORE_ALIAS}]'
|
||||
return await self._run_knn_search(client, index, query, query_embedding, k)
|
||||
|
||||
@staticmethod
|
||||
def _wrap_pre(pre: str) -> str:
|
||||
"""Parenthesize a multi-condition pre-filter before the ``=>`` KNN clause.
|
||||
|
||||
When ``pre`` combines several terms (e.g. ``@file_id:{x} @document:term``)
|
||||
the Valkey Search parser can otherwise mis-associate only the last term
|
||||
with the KNN clause. Wrapping the whole expression forces correct
|
||||
grouping. A bare ``*`` (match-all) and single-term expressions are left
|
||||
untouched.
|
||||
"""
|
||||
if pre and pre != '*' and ' ' in pre.strip():
|
||||
return f'({pre})'
|
||||
return pre
|
||||
|
||||
async def _run_knn_search(
|
||||
self,
|
||||
client: GlideClient,
|
||||
index: str,
|
||||
query: str,
|
||||
query_embedding: list[float],
|
||||
k: int,
|
||||
) -> dict[str, Any]:
|
||||
options = FtSearchOptions(
|
||||
params={'BLOB': self._pack_vector(list(query_embedding))},
|
||||
return_fields=[
|
||||
ReturnField(field_identifier=_VEC_SCORE_ALIAS, alias='distance'),
|
||||
ReturnField(field_identifier=_FIELD_DOCUMENT),
|
||||
ReturnField(field_identifier=_FIELD_METADATA),
|
||||
],
|
||||
limit=FtSearchLimit(0, k),
|
||||
dialect=2,
|
||||
)
|
||||
try:
|
||||
reply = await ft.search(client, index, query, options)
|
||||
except Exception as exc:
|
||||
if self._is_missing_index_error(exc):
|
||||
return {'ids': [[]], 'metadatas': [[]], 'distances': [[]]}
|
||||
raise
|
||||
return self._reply_to_chroma(index, reply, has_distance=True)
|
||||
|
||||
async def _run_text_search(
|
||||
self,
|
||||
client: GlideClient,
|
||||
index: str,
|
||||
query: str,
|
||||
k: int,
|
||||
) -> dict[str, Any]:
|
||||
options = FtSearchOptions(
|
||||
return_fields=[
|
||||
ReturnField(field_identifier=_FIELD_DOCUMENT),
|
||||
ReturnField(field_identifier=_FIELD_METADATA),
|
||||
],
|
||||
limit=FtSearchLimit(0, k),
|
||||
dialect=2,
|
||||
)
|
||||
try:
|
||||
reply = await ft.search(client, index, query, options)
|
||||
except Exception as exc:
|
||||
if self._is_missing_index_error(exc):
|
||||
return {'ids': [[]], 'metadatas': [[]], 'distances': [[]]}
|
||||
raise
|
||||
return self._reply_to_chroma(index, reply, has_distance=False)
|
||||
|
||||
@staticmethod
|
||||
def _is_missing_index_error(exc: Exception) -> bool:
|
||||
"""Return True if *exc* indicates the FT index does not exist.
|
||||
|
||||
``FT.DROPINDEX`` is applied eventually, so an index can briefly still
|
||||
appear in ``FT._LIST`` after being dropped; a follow-up search then
|
||||
fails with a "not found" error which we treat as an empty result.
|
||||
"""
|
||||
message = str(exc).lower()
|
||||
return 'not found' in message and 'index' in message
|
||||
|
||||
def _iter_reply_docs(self, reply: Any, prefix: str):
|
||||
"""Yield ``(doc_id, decoded_fields)`` pairs from an FT.SEARCH reply.
|
||||
|
||||
glide returns ``[total, {key: {field: value}, ...}]``. This shared
|
||||
iterator decodes each key, strips the per-collection prefix to recover
|
||||
the original document id, and decodes the field map — the logic both
|
||||
``_reply_to_chroma`` and ``list_by_filter`` need.
|
||||
"""
|
||||
docs = reply[1] if reply and len(reply) >= 2 and isinstance(reply[1], dict) else {}
|
||||
for key, fields in docs.items():
|
||||
key_str = self._decode(key)
|
||||
doc_id = key_str[len(prefix) :] if prefix and key_str.startswith(prefix) else key_str
|
||||
decoded_fields = {self._decode(fk): fv for fk, fv in fields.items()} if isinstance(fields, dict) else {}
|
||||
yield doc_id, decoded_fields
|
||||
|
||||
def _reply_to_chroma(self, index: str, reply: Any, has_distance: bool) -> dict[str, Any]:
|
||||
"""Convert an FT.SEARCH reply into Chroma-style nested lists.
|
||||
|
||||
The KNN score field (aliased ``distance``) is a COSINE/L2 distance
|
||||
directly, so no inversion is needed (unlike Qdrant).
|
||||
"""
|
||||
ids: list[str] = []
|
||||
distances: list[float] = []
|
||||
metadatas: list[dict[str, Any]] = []
|
||||
|
||||
if not reply or len(reply) < 2:
|
||||
return {'ids': [ids], 'metadatas': [metadatas], 'distances': [distances]}
|
||||
|
||||
prefix = self._key_prefix(index[len('idx:') :]) if index.startswith('idx:') else ''
|
||||
|
||||
for doc_id, decoded_fields in self._iter_reply_docs(reply, prefix):
|
||||
ids.append(doc_id)
|
||||
|
||||
if has_distance and 'distance' in decoded_fields:
|
||||
try:
|
||||
distances.append(float(self._decode(decoded_fields['distance'])))
|
||||
except (TypeError, ValueError):
|
||||
distances.append(0.0)
|
||||
else:
|
||||
distances.append(0.0)
|
||||
|
||||
metadata: dict[str, Any] = {}
|
||||
raw_meta = decoded_fields.get(_FIELD_METADATA)
|
||||
if raw_meta is not None:
|
||||
try:
|
||||
metadata = json.loads(self._decode(raw_meta))
|
||||
except (TypeError, ValueError):
|
||||
metadata = {}
|
||||
metadatas.append(metadata)
|
||||
|
||||
return {'ids': [ids], 'metadatas': [metadatas], 'distances': [distances]}
|
||||
|
||||
async def delete_by_file_id(self, collection: str, file_id: str) -> None:
|
||||
client = await self._ensure_client()
|
||||
index = self._index_name(collection)
|
||||
if not await self._index_exists(client, index):
|
||||
self.ap.logger.warning(f"Valkey Search collection '{collection}' not found for deletion")
|
||||
return
|
||||
|
||||
query = f'@{_FIELD_FILE_ID}:{{{self._encode_and_escape_tag(file_id)}}}'
|
||||
keys = await self._search_keys(client, index, query)
|
||||
if keys:
|
||||
await client.delete(keys)
|
||||
self.ap.logger.info(
|
||||
f"Deleted {len(keys)} embeddings from Valkey Search collection '{collection}' with file_id: {file_id}"
|
||||
)
|
||||
|
||||
async def delete_by_filter(self, collection: str, filter: dict[str, Any]) -> int:
|
||||
client = await self._ensure_client()
|
||||
index = self._index_name(collection)
|
||||
if not await self._index_exists(client, index):
|
||||
self.ap.logger.warning(f"Valkey Search collection '{collection}' not found for deletion")
|
||||
return 0
|
||||
|
||||
# Guard against accidental mass deletion: a non-empty filter that maps
|
||||
# to no usable (indexed) conditions must NOT fall back to match-all and
|
||||
# wipe the whole collection. Skip instead (matching Milvus / pgvector).
|
||||
query = self._triples_to_ft(filter)
|
||||
if not query:
|
||||
self.ap.logger.warning(
|
||||
"Valkey Search delete_by_filter on '%s': filter produced no usable conditions, skipping",
|
||||
collection,
|
||||
)
|
||||
return 0
|
||||
keys = await self._search_keys(client, index, query)
|
||||
if keys:
|
||||
await client.delete(keys)
|
||||
self.ap.logger.info(f"Deleted {len(keys)} embeddings from Valkey Search collection '{collection}' by filter")
|
||||
return len(keys)
|
||||
|
||||
async def list_by_filter(
|
||||
self,
|
||||
collection: str,
|
||||
filter: dict[str, Any] | None = None,
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
) -> tuple[list[dict[str, Any]], int]:
|
||||
client = await self._ensure_client()
|
||||
index = self._index_name(collection)
|
||||
if not await self._index_exists(client, index):
|
||||
return [], 0
|
||||
|
||||
query = self._triples_to_ft(filter) or _MATCH_ALL
|
||||
options = FtSearchOptions(
|
||||
return_fields=[
|
||||
ReturnField(field_identifier=_FIELD_DOCUMENT),
|
||||
ReturnField(field_identifier=_FIELD_METADATA),
|
||||
],
|
||||
limit=FtSearchLimit(offset, limit),
|
||||
dialect=2,
|
||||
)
|
||||
try:
|
||||
reply = await ft.search(client, index, query, options)
|
||||
except Exception as exc:
|
||||
if self._is_missing_index_error(exc):
|
||||
return [], 0
|
||||
raise
|
||||
|
||||
total = 0
|
||||
if reply:
|
||||
try:
|
||||
total = int(reply[0])
|
||||
except (TypeError, ValueError):
|
||||
total = 0
|
||||
|
||||
prefix = self._key_prefix(collection)
|
||||
items: list[dict[str, Any]] = []
|
||||
for doc_id, decoded_fields in self._iter_reply_docs(reply, prefix):
|
||||
document = decoded_fields.get(_FIELD_DOCUMENT)
|
||||
metadata: dict[str, Any] = {}
|
||||
raw_meta = decoded_fields.get(_FIELD_METADATA)
|
||||
if raw_meta is not None:
|
||||
try:
|
||||
metadata = json.loads(self._decode(raw_meta))
|
||||
except (TypeError, ValueError):
|
||||
metadata = {}
|
||||
|
||||
items.append(
|
||||
{
|
||||
'id': doc_id,
|
||||
'document': self._decode(document) if document is not None else None,
|
||||
'metadata': metadata,
|
||||
}
|
||||
)
|
||||
|
||||
return items, total
|
||||
|
||||
async def delete_collection(self, collection: str):
|
||||
client = await self._ensure_client()
|
||||
index = self._index_name(collection)
|
||||
self._ensured_indexes.discard(index)
|
||||
|
||||
if await self._index_exists(client, index):
|
||||
try:
|
||||
await ft.dropindex(client, index)
|
||||
except RequestError:
|
||||
# The index was already dropped (e.g. by a concurrent process)
|
||||
# between the existence check and this call — benign. Other
|
||||
# errors (connection / auth) must propagate so the caller knows
|
||||
# the operation failed rather than silently SCAN-deleting next.
|
||||
pass
|
||||
|
||||
# DROPINDEX does not remove the underlying hashes; delete them too.
|
||||
prefix = self._key_prefix(collection)
|
||||
cursor = b'0'
|
||||
deleted = 0
|
||||
for _ in range(_MAX_SCAN_ROUNDS):
|
||||
cursor, keys = await client.scan(cursor, match=f'{prefix}*', count=500)
|
||||
if keys:
|
||||
await client.delete(keys)
|
||||
deleted += len(keys)
|
||||
if cursor in (b'0', '0', 0):
|
||||
break
|
||||
self.ap.logger.info(f"Valkey Search collection '{collection}' deleted ({deleted} keys removed)")
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Internal search helpers
|
||||
# ------------------------------------------------------------------ #
|
||||
async def _index_exists(self, client: GlideClient, index: str) -> bool:
|
||||
if index in self._ensured_indexes:
|
||||
return True
|
||||
# ft.info is O(1) and raises RequestError when the index does not
|
||||
# exist, vs ft.list which is O(n) over every index on the server and
|
||||
# was being paid on the first query to each collection.
|
||||
try:
|
||||
await ft.info(client, index)
|
||||
self._ensured_indexes.add(index)
|
||||
return True
|
||||
except RequestError:
|
||||
return False
|
||||
|
||||
async def _search_keys(self, client: GlideClient, index: str, query: str) -> list[str]:
|
||||
"""Return all matching document keys for a query (NOCONTENT).
|
||||
|
||||
Paginates through the full result set in pages of ``_DELETE_SCAN_BATCH``
|
||||
so that queries matching more than one page of chunks are fully
|
||||
enumerated (avoids silently truncating deletes and leaving orphaned
|
||||
vectors).
|
||||
"""
|
||||
keys: list[str] = []
|
||||
offset = 0
|
||||
while True:
|
||||
options = FtSearchOptions(
|
||||
nocontent=True,
|
||||
limit=FtSearchLimit(offset, _DELETE_SCAN_BATCH),
|
||||
dialect=2,
|
||||
)
|
||||
try:
|
||||
reply = await ft.search(client, index, query, options)
|
||||
except Exception as exc:
|
||||
if self._is_missing_index_error(exc):
|
||||
return keys
|
||||
raise
|
||||
|
||||
if not reply or len(reply) < 2:
|
||||
break
|
||||
|
||||
# reply[0] is the total match count; reply[1] holds this page.
|
||||
total = 0
|
||||
try:
|
||||
total = int(reply[0])
|
||||
except (TypeError, ValueError):
|
||||
total = 0
|
||||
|
||||
docs = reply[1]
|
||||
if isinstance(docs, dict):
|
||||
page = [self._decode(k) for k in docs.keys()]
|
||||
elif isinstance(docs, (list, tuple)):
|
||||
page = [self._decode(k) for k in docs]
|
||||
else:
|
||||
page = []
|
||||
|
||||
if not page:
|
||||
break
|
||||
keys.extend(page)
|
||||
|
||||
offset += len(page)
|
||||
if offset >= total or len(page) < _DELETE_SCAN_BATCH:
|
||||
break
|
||||
|
||||
return keys
|
||||
@@ -0,0 +1,204 @@
|
||||
"""Workflow-Pipeline通信适配器
|
||||
|
||||
这个模块提供了Workflow和Pipeline之间的通信适配,使用SDK标准的MessageEnvelope格式。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _WorkflowPipelineCaptureAdapter:
|
||||
"""Workflow-Pipeline通信适配器
|
||||
|
||||
用于在Workflow节点和Pipeline之间进行标准化的消息传递。
|
||||
支持MessageEnvelope格式的双向转换。
|
||||
"""
|
||||
|
||||
def __init__(self, context: Any):
|
||||
"""初始化适配器
|
||||
|
||||
Args:
|
||||
context: ExecutionContext - Workflow执行上下文
|
||||
"""
|
||||
self.context = context
|
||||
self.responses: list[dict[str, Any]] = []
|
||||
self.bot_account_id: Optional[str] = None
|
||||
self._logger = logging.getLogger(__name__)
|
||||
|
||||
async def call_pipeline_with_envelope(
|
||||
self,
|
||||
envelope: Any,
|
||||
pipeline_executor: Any
|
||||
) -> Any:
|
||||
"""使用MessageEnvelope调用Pipeline
|
||||
|
||||
Args:
|
||||
envelope: MessageEnvelope - 标准消息信封
|
||||
pipeline_executor: Pipeline执行器实例
|
||||
|
||||
Returns:
|
||||
MessageEnvelope - 执行结果信封
|
||||
"""
|
||||
try:
|
||||
# 动态导入以避免循环依赖
|
||||
from langbot_plugin_sdk.workflow import envelope_to_query, query_to_envelope
|
||||
|
||||
# 1. 转换为Query
|
||||
query = envelope_to_query(envelope)
|
||||
|
||||
# 2. 调用Pipeline
|
||||
result_query = await pipeline_executor.execute(query)
|
||||
|
||||
# 3. 转换回Envelope
|
||||
result_envelope = query_to_envelope(result_query, envelope)
|
||||
|
||||
self._logger.debug(
|
||||
f'Pipeline execution completed for workflow {envelope.workflow_id}',
|
||||
extra={
|
||||
'workflow_id': envelope.workflow_id,
|
||||
'execution_id': envelope.execution_id,
|
||||
'node_id': envelope.node_id,
|
||||
}
|
||||
)
|
||||
|
||||
return result_envelope
|
||||
|
||||
except Exception as e:
|
||||
self._logger.error(
|
||||
f'Pipeline execution failed: {e}',
|
||||
exc_info=True,
|
||||
extra={
|
||||
'workflow_id': envelope.workflow_id,
|
||||
'execution_id': envelope.execution_id,
|
||||
'node_id': envelope.node_id,
|
||||
}
|
||||
)
|
||||
raise
|
||||
|
||||
def validate_envelope(self, envelope: Any) -> bool:
|
||||
"""验证MessageEnvelope的有效性
|
||||
|
||||
Args:
|
||||
envelope: MessageEnvelope - 要验证的消息信封
|
||||
|
||||
Returns:
|
||||
bool - 验证是否通过
|
||||
"""
|
||||
required_fields = [
|
||||
'message_id',
|
||||
'workflow_id',
|
||||
'node_id',
|
||||
'execution_id',
|
||||
'payload',
|
||||
'launcher_type',
|
||||
]
|
||||
|
||||
for field in required_fields:
|
||||
if not hasattr(envelope, field):
|
||||
self._logger.warning(
|
||||
f'MessageEnvelope missing required field: {field}'
|
||||
)
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def get_responses(self) -> list[dict[str, Any]]:
|
||||
"""获取所有响应
|
||||
|
||||
Returns:
|
||||
list - 响应列表
|
||||
"""
|
||||
return self.responses.copy()
|
||||
|
||||
def add_response(self, response: dict[str, Any]) -> None:
|
||||
"""添加响应
|
||||
|
||||
Args:
|
||||
response: dict - 响应数据
|
||||
"""
|
||||
self.responses.append(response)
|
||||
|
||||
def get_last_text_response(self) -> str:
|
||||
"""获取最后一个文本响应
|
||||
|
||||
Returns:
|
||||
str - 最后一个响应的文本内容
|
||||
"""
|
||||
if not self.responses:
|
||||
return ''
|
||||
|
||||
last_response = self.responses[-1]
|
||||
return str(last_response.get('content', '') or '')
|
||||
|
||||
def clear_responses(self) -> None:
|
||||
"""清空所有响应"""
|
||||
self.responses.clear()
|
||||
|
||||
|
||||
class WorkflowPipelineCompatibilityLayer:
|
||||
"""Workflow-Pipeline兼容性层
|
||||
|
||||
提供向后兼容性,支持旧的Pipeline Query格式和新的MessageEnvelope格式。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""初始化兼容性层"""
|
||||
self._logger = logging.getLogger(__name__)
|
||||
|
||||
def is_workflow_context(self, query: Any) -> bool:
|
||||
"""检查Query是否包含Workflow上下文
|
||||
|
||||
Args:
|
||||
query: Query - Pipeline Query对象
|
||||
|
||||
Returns:
|
||||
bool - 是否来自Workflow
|
||||
"""
|
||||
if hasattr(query, 'is_from_workflow'):
|
||||
return query.is_from_workflow()
|
||||
|
||||
if hasattr(query, 'get_workflow_context'):
|
||||
context = query.get_workflow_context()
|
||||
return bool(context and context.get('workflow_id'))
|
||||
|
||||
return False
|
||||
|
||||
def get_workflow_id(self, query: Any) -> Optional[str]:
|
||||
"""从Query获取Workflow ID
|
||||
|
||||
Args:
|
||||
query: Query - Pipeline Query对象
|
||||
|
||||
Returns:
|
||||
str - Workflow ID,如果不存在则返回None
|
||||
"""
|
||||
if hasattr(query, 'get_workflow_id'):
|
||||
return query.get_workflow_id()
|
||||
|
||||
if hasattr(query, 'get_workflow_context'):
|
||||
context = query.get_workflow_context()
|
||||
return context.get('workflow_id') if context else None
|
||||
|
||||
return None
|
||||
|
||||
def get_execution_id(self, query: Any) -> Optional[str]:
|
||||
"""从Query获取执行ID
|
||||
|
||||
Args:
|
||||
query: Query - Pipeline Query对象
|
||||
|
||||
Returns:
|
||||
str - 执行ID,如果不存在则返回None
|
||||
"""
|
||||
if hasattr(query, 'get_execution_id'):
|
||||
return query.get_execution_id()
|
||||
|
||||
if hasattr(query, 'get_workflow_context'):
|
||||
context = query.get_workflow_context()
|
||||
return context.get('execution_id') if context else None
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,504 @@
|
||||
"""Workflow debug execution support.
|
||||
|
||||
This module provides debugging capabilities for workflow execution, including:
|
||||
- ExecutionLog: Structured log entries for execution tracking
|
||||
- DebugExecutionState: State management for debug sessions (pause, resume, breakpoints)
|
||||
- DebugWorkflowExecutor: Extended executor with step-by-step debugging support
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import traceback
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any, Optional, TYPE_CHECKING
|
||||
|
||||
from .entities import (
|
||||
WorkflowDefinition,
|
||||
NodeDefinition,
|
||||
EdgeDefinition,
|
||||
ExecutionContext,
|
||||
ExecutionStatus,
|
||||
NodeState,
|
||||
NodeStatus,
|
||||
)
|
||||
from .executor import WorkflowExecutor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core import app
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ExecutionLog:
|
||||
"""Execution log entry"""
|
||||
|
||||
def __init__(self, level: str, message: str, node_id: Optional[str] = None, data: Optional[dict] = None):
|
||||
self.id = str(uuid.uuid4())
|
||||
self.timestamp = datetime.now().isoformat()
|
||||
self.level = level
|
||||
self.message = message
|
||||
self.node_id = node_id
|
||||
self.data = data or {}
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
'id': self.id,
|
||||
'timestamp': self.timestamp,
|
||||
'level': self.level,
|
||||
'message': self.message,
|
||||
'node_id': self.node_id,
|
||||
'data': self.data,
|
||||
}
|
||||
|
||||
|
||||
class DebugExecutionState:
|
||||
"""State for a debug execution"""
|
||||
|
||||
def __init__(self, execution_id: str, breakpoints: list[str] = None):
|
||||
self.execution_id = execution_id
|
||||
self.status: str = 'running'
|
||||
self.is_paused: bool = False
|
||||
self.is_stopped: bool = False
|
||||
self.current_node_id: Optional[str] = None
|
||||
self.breakpoints: set[str] = set(breakpoints or [])
|
||||
self.logs: list[ExecutionLog] = []
|
||||
self.pending_logs: list[ExecutionLog] = []
|
||||
self._pause_event = asyncio.Event()
|
||||
self._pause_event.set() # Initially not paused
|
||||
self._stop_event = asyncio.Event()
|
||||
|
||||
def add_log(self, level: str, message: str, node_id: str = None, data: dict = None):
|
||||
"""Add a log entry"""
|
||||
log = ExecutionLog(level, message, node_id, data)
|
||||
self.logs.append(log)
|
||||
self.pending_logs.append(log)
|
||||
logger.log(
|
||||
getattr(logging, level.upper(), logging.INFO),
|
||||
f'[Workflow Debug] {message}',
|
||||
extra={'node_id': node_id, 'data': data},
|
||||
)
|
||||
|
||||
def get_pending_logs(self) -> list[dict]:
|
||||
"""Get and clear pending logs"""
|
||||
logs = [log.to_dict() for log in self.pending_logs]
|
||||
self.pending_logs = []
|
||||
return logs
|
||||
|
||||
def pause(self):
|
||||
"""Pause execution"""
|
||||
self.is_paused = True
|
||||
self._pause_event.clear()
|
||||
self.add_log('info', 'Execution paused')
|
||||
|
||||
def resume(self):
|
||||
"""Resume execution"""
|
||||
self.is_paused = False
|
||||
self._pause_event.set()
|
||||
self.add_log('info', 'Execution resumed')
|
||||
|
||||
def stop(self):
|
||||
"""Stop execution"""
|
||||
self.is_stopped = True
|
||||
self.status = 'cancelled'
|
||||
self._stop_event.set()
|
||||
self._pause_event.set() # Release any pause
|
||||
self.add_log('info', 'Execution stopped')
|
||||
|
||||
async def wait_if_paused(self):
|
||||
"""Wait if execution is paused"""
|
||||
if self.is_paused:
|
||||
self.add_log('info', 'Waiting for resume...')
|
||||
await self._pause_event.wait()
|
||||
|
||||
def check_breakpoint(self, node_id: str) -> bool:
|
||||
"""Check if there's a breakpoint at the given node"""
|
||||
return node_id in self.breakpoints
|
||||
|
||||
|
||||
class DebugWorkflowExecutor(WorkflowExecutor):
|
||||
"""
|
||||
Debug-enabled workflow executor with step-by-step execution support.
|
||||
Extends WorkflowExecutor with debugging capabilities.
|
||||
"""
|
||||
|
||||
# Class-level storage for active debug sessions
|
||||
_debug_states: dict[str, DebugExecutionState] = {}
|
||||
|
||||
def __init__(self, ap: Optional['app.Application'] = None):
|
||||
super().__init__(ap)
|
||||
|
||||
@classmethod
|
||||
def get_debug_state(cls, execution_id: str) -> Optional[DebugExecutionState]:
|
||||
"""Get debug state for an execution"""
|
||||
return cls._debug_states.get(execution_id)
|
||||
|
||||
@classmethod
|
||||
def create_debug_state(cls, execution_id: str, breakpoints: list[str] = None) -> DebugExecutionState:
|
||||
"""Create a new debug state"""
|
||||
state = DebugExecutionState(execution_id, breakpoints)
|
||||
cls._debug_states[execution_id] = state
|
||||
return state
|
||||
|
||||
@classmethod
|
||||
def remove_debug_state(cls, execution_id: str):
|
||||
"""Remove debug state for an execution"""
|
||||
cls._debug_states.pop(execution_id, None)
|
||||
|
||||
async def execute_debug(
|
||||
self,
|
||||
workflow: WorkflowDefinition,
|
||||
context: ExecutionContext,
|
||||
debug_state: DebugExecutionState,
|
||||
) -> ExecutionContext:
|
||||
"""
|
||||
Execute a workflow in debug mode.
|
||||
|
||||
Args:
|
||||
workflow: Workflow definition
|
||||
context: Execution context
|
||||
debug_state: Debug execution state
|
||||
|
||||
Returns:
|
||||
Updated execution context
|
||||
"""
|
||||
context.status = ExecutionStatus.RUNNING
|
||||
context.start_time = datetime.now()
|
||||
debug_state.add_log('info', f'Starting debug execution for workflow: {workflow.name}')
|
||||
|
||||
try:
|
||||
# Build execution graph
|
||||
node_map = {node.id: node for node in workflow.nodes}
|
||||
edge_map = self._build_edge_map(workflow.edges)
|
||||
self._edges = workflow.edges
|
||||
|
||||
# Initialize node states
|
||||
for node in workflow.nodes:
|
||||
if node.id not in context.node_states:
|
||||
context.node_states[node.id] = NodeState(node_id=node.id)
|
||||
|
||||
# Find start node(s)
|
||||
start_nodes = self._find_start_nodes(workflow.nodes, workflow.edges)
|
||||
|
||||
if not start_nodes:
|
||||
raise ValueError('No start nodes found in workflow')
|
||||
|
||||
debug_state.add_log('info', f'Found {len(start_nodes)} start node(s)')
|
||||
|
||||
# Execute from start nodes
|
||||
for start_node in start_nodes:
|
||||
if debug_state.is_stopped:
|
||||
break
|
||||
|
||||
await self._execute_debug_from_node(
|
||||
start_node, node_map, edge_map, context, debug_state, workflow.settings.max_retries
|
||||
)
|
||||
|
||||
# Set final status
|
||||
if debug_state.is_stopped:
|
||||
context.status = ExecutionStatus.CANCELLED
|
||||
debug_state.status = 'cancelled'
|
||||
else:
|
||||
all_completed = all(
|
||||
state.status in (NodeStatus.COMPLETED, NodeStatus.SKIPPED) for state in context.node_states.values()
|
||||
)
|
||||
|
||||
if all_completed:
|
||||
context.status = ExecutionStatus.COMPLETED
|
||||
debug_state.status = 'completed'
|
||||
debug_state.add_log('info', 'Workflow execution completed successfully')
|
||||
else:
|
||||
has_failed = any(state.status == NodeStatus.FAILED for state in context.node_states.values())
|
||||
if has_failed:
|
||||
context.status = ExecutionStatus.FAILED
|
||||
debug_state.status = 'error'
|
||||
|
||||
except Exception as e:
|
||||
context.status = ExecutionStatus.FAILED
|
||||
context.error = str(e)
|
||||
debug_state.status = 'error'
|
||||
debug_state.add_log('error', f'Workflow execution failed: {e}', data={'traceback': traceback.format_exc()})
|
||||
logger.error(f'Debug workflow execution failed: {e}\n{traceback.format_exc()}')
|
||||
|
||||
finally:
|
||||
context.end_time = datetime.now()
|
||||
|
||||
return context
|
||||
|
||||
async def _execute_debug_from_node(
|
||||
self,
|
||||
node: NodeDefinition,
|
||||
node_map: dict[str, NodeDefinition],
|
||||
edge_map: dict[str, list[EdgeDefinition]],
|
||||
context: ExecutionContext,
|
||||
debug_state: DebugExecutionState,
|
||||
max_retries: int = 3,
|
||||
):
|
||||
"""Execute workflow from a node with debug support"""
|
||||
|
||||
# Check if stopped
|
||||
if debug_state.is_stopped:
|
||||
return
|
||||
|
||||
# Wait if paused
|
||||
await debug_state.wait_if_paused()
|
||||
|
||||
# Check if should skip
|
||||
if await self._should_skip_node(node, context):
|
||||
if context.node_states[node.id].status == NodeStatus.SKIPPED:
|
||||
debug_state.add_log('info', f'Skipping node: {node.id}', node_id=node.id)
|
||||
return
|
||||
|
||||
# Check breakpoint
|
||||
if debug_state.check_breakpoint(node.id):
|
||||
debug_state.add_log('info', f'Hit breakpoint at node: {node.id}', node_id=node.id)
|
||||
debug_state.pause()
|
||||
await debug_state.wait_if_paused()
|
||||
|
||||
# Update current node
|
||||
debug_state.current_node_id = node.id
|
||||
debug_state.add_log('info', f'Executing node: {node.id} ({node.type})', node_id=node.id)
|
||||
|
||||
# Execute node
|
||||
await self._execute_debug_node(node, context, debug_state, max_retries)
|
||||
|
||||
# Check if stopped or failed
|
||||
if debug_state.is_stopped:
|
||||
return
|
||||
if context.node_states[node.id].status == NodeStatus.FAILED:
|
||||
return
|
||||
|
||||
# Get outgoing edges
|
||||
outgoing_edges = edge_map.get(node.id, [])
|
||||
|
||||
# Execute next nodes
|
||||
for edge in outgoing_edges:
|
||||
if debug_state.is_stopped:
|
||||
break
|
||||
|
||||
target_node = node_map.get(edge.target_node)
|
||||
if not target_node:
|
||||
continue
|
||||
|
||||
# Check edge condition
|
||||
if edge.condition:
|
||||
condition_met = await self._evaluate_condition(edge.condition, context)
|
||||
if not condition_met:
|
||||
debug_state.add_log('debug', f'Edge condition not met: {edge.condition}', node_id=node.id)
|
||||
continue
|
||||
|
||||
# Check if all inputs are ready
|
||||
if await self._inputs_ready(target_node, edge_map, context):
|
||||
await self._execute_debug_from_node(target_node, node_map, edge_map, context, debug_state, max_retries)
|
||||
|
||||
async def _execute_debug_node(
|
||||
self, node: NodeDefinition, context: ExecutionContext, debug_state: DebugExecutionState, max_retries: int = 3
|
||||
):
|
||||
"""Execute a single node with debug logging"""
|
||||
|
||||
node_state = context.node_states[node.id]
|
||||
node_state.status = NodeStatus.RUNNING
|
||||
node_state.start_time = datetime.now()
|
||||
|
||||
# Get node instance (pass ap for access to services)
|
||||
node_instance = self.registry.create_instance(node.type, node.id, node.config, ap=self.ap)
|
||||
|
||||
if not node_instance:
|
||||
node_state.status = NodeStatus.FAILED
|
||||
node_state.error = f'Unknown node type: {node.type}'
|
||||
node_state.end_time = datetime.now()
|
||||
debug_state.add_log('error', f'Unknown node type: {node.type}', node_id=node.id)
|
||||
self._record_execution_step(node, node_state, context)
|
||||
await self._persist_node_execution(node, node_state, context)
|
||||
return
|
||||
|
||||
# Resolve inputs
|
||||
inputs = await self._resolve_inputs(node, context)
|
||||
node_state.inputs = inputs
|
||||
debug_state.add_log(
|
||||
'debug', 'Node inputs resolved', node_id=node.id, data={'inputs': self._safe_serialize(inputs)}
|
||||
)
|
||||
|
||||
# Validate inputs
|
||||
validation_errors = await node_instance.validate_inputs(inputs)
|
||||
if validation_errors:
|
||||
node_state.status = NodeStatus.FAILED
|
||||
node_state.error = '; '.join(validation_errors)
|
||||
node_state.end_time = datetime.now()
|
||||
debug_state.add_log('error', f'Input validation failed: {node_state.error}', node_id=node.id)
|
||||
self._record_execution_step(node, node_state, context)
|
||||
await self._persist_node_execution(node, node_state, context)
|
||||
return
|
||||
|
||||
# Execute with retries
|
||||
for attempt in range(max_retries + 1):
|
||||
if debug_state.is_stopped:
|
||||
node_state.status = NodeStatus.FAILED
|
||||
node_state.error = 'Execution stopped'
|
||||
node_state.end_time = datetime.now()
|
||||
break
|
||||
|
||||
try:
|
||||
outputs = await node_instance.execute(inputs, context)
|
||||
node_state.outputs = outputs
|
||||
node_state.status = NodeStatus.COMPLETED
|
||||
node_state.end_time = datetime.now()
|
||||
|
||||
duration_ms = int((node_state.end_time - node_state.start_time).total_seconds() * 1000)
|
||||
debug_state.add_log(
|
||||
'info',
|
||||
f'Node completed in {duration_ms}ms',
|
||||
node_id=node.id,
|
||||
data={'outputs': self._safe_serialize(outputs), 'duration_ms': duration_ms},
|
||||
)
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
node_state.retry_count = attempt + 1
|
||||
debug_state.add_log(
|
||||
'warning', f'Node execution failed (attempt {attempt + 1}/{max_retries + 1}): {e}', node_id=node.id
|
||||
)
|
||||
|
||||
if attempt < max_retries:
|
||||
await asyncio.sleep(1)
|
||||
else:
|
||||
node_state.status = NodeStatus.FAILED
|
||||
node_state.error = str(e)
|
||||
node_state.end_time = datetime.now()
|
||||
debug_state.add_log(
|
||||
'error',
|
||||
f'Node failed after {max_retries + 1} attempts: {e}',
|
||||
node_id=node.id,
|
||||
data={'error': str(e), 'traceback': traceback.format_exc()},
|
||||
)
|
||||
|
||||
self._record_execution_step(node, node_state, context)
|
||||
await self._persist_node_execution(node, node_state, context)
|
||||
|
||||
async def step_execute(
|
||||
self,
|
||||
workflow: WorkflowDefinition,
|
||||
context: ExecutionContext,
|
||||
debug_state: DebugExecutionState,
|
||||
) -> dict:
|
||||
"""
|
||||
Execute one step (one node) in debug mode.
|
||||
|
||||
Returns:
|
||||
Dict with node_id, node_state, and completed status
|
||||
"""
|
||||
# Find next node to execute
|
||||
next_node = self._find_next_executable_node(workflow, context)
|
||||
|
||||
if not next_node:
|
||||
debug_state.status = 'completed'
|
||||
return {'completed': True}
|
||||
|
||||
# Execute single node
|
||||
debug_state.current_node_id = next_node.id
|
||||
await self._execute_debug_node(next_node, context, debug_state, workflow.settings.max_retries)
|
||||
|
||||
node_state = context.node_states.get(next_node.id)
|
||||
|
||||
# Check if workflow is complete
|
||||
all_done = all(
|
||||
state.status in (NodeStatus.COMPLETED, NodeStatus.SKIPPED, NodeStatus.FAILED)
|
||||
for state in context.node_states.values()
|
||||
)
|
||||
|
||||
if all_done:
|
||||
debug_state.status = 'completed'
|
||||
context.status = ExecutionStatus.COMPLETED
|
||||
|
||||
return {
|
||||
'node_id': next_node.id,
|
||||
'node_state': {
|
||||
'status': node_state.status.value if node_state else 'unknown',
|
||||
'inputs': self._safe_serialize(node_state.inputs) if node_state else {},
|
||||
'outputs': self._safe_serialize(node_state.outputs) if node_state else {},
|
||||
'error': node_state.error if node_state else None,
|
||||
},
|
||||
'completed': all_done,
|
||||
}
|
||||
|
||||
def _find_next_executable_node(
|
||||
self, workflow: WorkflowDefinition, context: ExecutionContext
|
||||
) -> Optional[NodeDefinition]:
|
||||
"""Find the next node that can be executed"""
|
||||
edge_map = self._build_edge_map(workflow.edges)
|
||||
|
||||
for node in workflow.nodes:
|
||||
state = context.node_states.get(node.id)
|
||||
|
||||
# Skip completed, running, or failed nodes
|
||||
if state and state.status in (
|
||||
NodeStatus.COMPLETED,
|
||||
NodeStatus.RUNNING,
|
||||
NodeStatus.FAILED,
|
||||
NodeStatus.SKIPPED,
|
||||
):
|
||||
continue
|
||||
|
||||
incoming_nodes = self._incoming_dependency_nodes(node.id, edge_map)
|
||||
|
||||
# If no incoming nodes, it's a start node
|
||||
if not incoming_nodes:
|
||||
return node
|
||||
|
||||
# Check if all incoming nodes are done
|
||||
all_incoming_done = True
|
||||
for source_id in incoming_nodes:
|
||||
source_state = context.node_states.get(source_id)
|
||||
if not source_state or source_state.status not in (NodeStatus.COMPLETED, NodeStatus.SKIPPED):
|
||||
all_incoming_done = False
|
||||
break
|
||||
|
||||
if all_incoming_done:
|
||||
return node
|
||||
|
||||
return None
|
||||
|
||||
def _safe_serialize(self, data: Any) -> Any:
|
||||
"""Safely serialize data for logging"""
|
||||
if data is None:
|
||||
return None
|
||||
if isinstance(data, (str, int, float, bool)):
|
||||
return data
|
||||
if isinstance(data, (list, tuple)):
|
||||
return [self._safe_serialize(item) for item in data[:100]] # Limit list size
|
||||
if isinstance(data, dict):
|
||||
result = {}
|
||||
for key, value in list(data.items())[:50]: # Limit dict size
|
||||
result[str(key)] = self._safe_serialize(value)
|
||||
return result
|
||||
# For complex objects, try to convert to string
|
||||
try:
|
||||
return str(data)[:1000] # Limit string length
|
||||
except Exception:
|
||||
return '<non-serializable>'
|
||||
|
||||
def get_execution_state(self, context: ExecutionContext, debug_state: DebugExecutionState) -> dict:
|
||||
"""Get current execution state for API response"""
|
||||
node_states = {}
|
||||
for node_id, state in context.node_states.items():
|
||||
node_states[node_id] = {
|
||||
'status': state.status.value,
|
||||
'inputs': self._safe_serialize(state.inputs),
|
||||
'outputs': self._safe_serialize(state.outputs),
|
||||
'error': state.error,
|
||||
'startTime': state.start_time.isoformat() if state.start_time else None,
|
||||
'endTime': state.end_time.isoformat() if state.end_time else None,
|
||||
'duration': int((state.end_time - state.start_time).total_seconds() * 1000)
|
||||
if state.start_time and state.end_time
|
||||
else None,
|
||||
}
|
||||
|
||||
return {
|
||||
'status': debug_state.status,
|
||||
'current_node_id': debug_state.current_node_id,
|
||||
'node_states': node_states,
|
||||
'new_logs': debug_state.get_pending_logs(),
|
||||
'error': context.error,
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
"""Workflow entities and data models
|
||||
|
||||
This module defines workflow entities using SDK standard entities where available,
|
||||
and local-specific entities for LangBot_copy-specific functionality.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any, Optional
|
||||
import pydantic
|
||||
|
||||
# Import SDK entities for standard workflow protocol types
|
||||
# These are re-exported for use by other modules in the workflow package.
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import (
|
||||
ExecutionContext as ExecutionContext,
|
||||
ExecutionStep as ExecutionStep,
|
||||
MessageContext as MessageContext,
|
||||
NodeDefinition,
|
||||
NodeState as NodeState,
|
||||
PortDefinition as PortDefinition,
|
||||
)
|
||||
from langbot_plugin.api.entities.builtin.workflow.enums import (
|
||||
ExecutionStatus as ExecutionStatus,
|
||||
NodeStatus as NodeStatus,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ExecutionContext",
|
||||
"ExecutionStep",
|
||||
"MessageContext",
|
||||
"NodeDefinition",
|
||||
"NodeState",
|
||||
"PortDefinition",
|
||||
"ExecutionStatus",
|
||||
"NodeStatus",
|
||||
]
|
||||
|
||||
|
||||
class Position(pydantic.BaseModel):
|
||||
"""Node position on canvas"""
|
||||
|
||||
x: float = 0
|
||||
y: float = 0
|
||||
|
||||
|
||||
class EdgeDefinition(pydantic.BaseModel):
|
||||
"""Workflow edge definition (connection between nodes)"""
|
||||
|
||||
id: str
|
||||
source_node: str
|
||||
source_port: str = 'output'
|
||||
target_node: str
|
||||
target_port: str = 'input'
|
||||
edge_type: str = 'legacy' # control, data, or legacy (old mixed semantics)
|
||||
condition: Optional[str] = None # Optional condition expression
|
||||
|
||||
|
||||
class TriggerDefinition(pydantic.BaseModel):
|
||||
"""Workflow trigger definition"""
|
||||
|
||||
id: str
|
||||
type: str # message, cron, event, webhook
|
||||
config: dict[str, Any] = {}
|
||||
enabled: bool = True
|
||||
|
||||
|
||||
class WorkflowSettings(pydantic.BaseModel):
|
||||
"""Workflow settings"""
|
||||
|
||||
# Execution settings
|
||||
max_execution_time: int = 300 # seconds
|
||||
max_retries: int = 3
|
||||
retry_delay: int = 5 # seconds
|
||||
|
||||
# Error handling
|
||||
error_handling: str = 'stop' # stop, continue, retry
|
||||
|
||||
# Logging
|
||||
log_level: str = 'info'
|
||||
save_execution_history: bool = True
|
||||
|
||||
# Concurrency
|
||||
max_concurrent_executions: int = 10
|
||||
|
||||
|
||||
class SafetyConfig(pydantic.BaseModel):
|
||||
"""Safety configuration (inherited from Pipeline)"""
|
||||
|
||||
content_filter: dict[str, Any] = {'enable': False, 'sensitive_words': [], 'replace_with': '***'}
|
||||
rate_limit: dict[str, Any] = {'enable': False, 'requests_per_minute': 60, 'burst_limit': 10}
|
||||
|
||||
|
||||
class OutputConfig(pydantic.BaseModel):
|
||||
"""Output configuration (inherited from Pipeline)"""
|
||||
|
||||
long_text_processing: dict[str, Any] = {
|
||||
'strategy': 'split', # split, truncate, file
|
||||
'max_length': 4000,
|
||||
'split_separator': '\n\n',
|
||||
}
|
||||
force_delay: dict[str, Any] = {'enable': False, 'min_delay_ms': 0, 'max_delay_ms': 0}
|
||||
misc: dict[str, Any] = {}
|
||||
|
||||
|
||||
class WorkflowGlobalConfig(pydantic.BaseModel):
|
||||
"""Workflow global configuration (inherited from Pipeline capabilities)"""
|
||||
|
||||
safety: SafetyConfig = SafetyConfig()
|
||||
output: OutputConfig = OutputConfig()
|
||||
|
||||
|
||||
class ExtensionsPreferences(pydantic.BaseModel):
|
||||
"""Extensions preferences (same as Pipeline)"""
|
||||
|
||||
enable_all_plugins: bool = True
|
||||
enable_all_mcp_servers: bool = True
|
||||
plugins: list[str] = []
|
||||
mcp_servers: list[str] = []
|
||||
|
||||
|
||||
class ConversationVariable(pydantic.BaseModel):
|
||||
"""Conversation-level variable definition"""
|
||||
|
||||
name: str
|
||||
type: str = 'string' # string, number, boolean, object, array
|
||||
description: str = ''
|
||||
default_value: Any = None
|
||||
max_length: Optional[int] = None # For strings
|
||||
|
||||
|
||||
class WorkflowDefinition(pydantic.BaseModel):
|
||||
"""Complete workflow definition"""
|
||||
|
||||
uuid: str
|
||||
name: str
|
||||
description: str = ''
|
||||
emoji: str = '💼'
|
||||
version: int = 1
|
||||
|
||||
# Workflow graph
|
||||
nodes: list[NodeDefinition] = []
|
||||
edges: list[EdgeDefinition] = []
|
||||
|
||||
# Variables
|
||||
variables: dict[str, Any] = {} # Global variables
|
||||
conversation_variables: list[ConversationVariable] = [] # Session-level variables
|
||||
|
||||
# Settings
|
||||
settings: WorkflowSettings = WorkflowSettings()
|
||||
|
||||
# Triggers (for automation)
|
||||
triggers: list[TriggerDefinition] = []
|
||||
|
||||
# Global configuration (inherited from Pipeline)
|
||||
global_config: WorkflowGlobalConfig = WorkflowGlobalConfig()
|
||||
|
||||
# Extensions
|
||||
extensions_preferences: ExtensionsPreferences = ExtensionsPreferences()
|
||||
|
||||
# Metadata
|
||||
is_enabled: bool = True
|
||||
created_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
|
||||
# Source tracking (for imported workflows)
|
||||
source: Optional[str] = None # dify, n8n, langflow, etc.
|
||||
source_id: Optional[str] = None
|
||||
@@ -0,0 +1,759 @@
|
||||
"""Workflow execution engine.
|
||||
|
||||
This module contains the core workflow execution logic:
|
||||
- WorkflowExecutor: Main execution engine with control flow handling
|
||||
- ParallelExecutor: Parallel branch execution
|
||||
- LoopExecutor: Loop/iterator execution
|
||||
|
||||
Debug execution support has been moved to the ``debug`` module.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import re
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any, Optional, TYPE_CHECKING
|
||||
|
||||
import sqlalchemy
|
||||
|
||||
from .entities import (
|
||||
WorkflowDefinition,
|
||||
NodeDefinition,
|
||||
EdgeDefinition,
|
||||
ExecutionContext,
|
||||
ExecutionStatus,
|
||||
NodeState,
|
||||
NodeStatus,
|
||||
ExecutionStep,
|
||||
)
|
||||
from ..entity.persistence import workflow as persistence_workflow
|
||||
from .registry import NodeTypeRegistry
|
||||
from .safe_eval import safe_eval_with_vars
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core import app
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class WorkflowExecutor:
|
||||
"""
|
||||
Workflow execution engine.
|
||||
Handles the execution of workflow definitions with proper control flow.
|
||||
"""
|
||||
|
||||
def __init__(self, ap: Optional['app.Application'] = None):
|
||||
self.ap = ap
|
||||
self.registry = NodeTypeRegistry.instance()
|
||||
self._edges: list[EdgeDefinition] = []
|
||||
|
||||
async def execute(
|
||||
self, workflow: WorkflowDefinition, context: ExecutionContext, start_node_id: Optional[str] = None
|
||||
) -> ExecutionContext:
|
||||
"""
|
||||
Execute a workflow.
|
||||
|
||||
Args:
|
||||
workflow: Workflow definition
|
||||
context: Execution context
|
||||
start_node_id: Optional starting node (for resumption)
|
||||
|
||||
Returns:
|
||||
Updated execution context
|
||||
"""
|
||||
context.status = ExecutionStatus.RUNNING
|
||||
context.start_time = datetime.now()
|
||||
|
||||
try:
|
||||
# Build execution graph
|
||||
node_map = {node.id: node for node in workflow.nodes}
|
||||
edge_map = self._build_edge_map(workflow.edges)
|
||||
self._edges = workflow.edges
|
||||
|
||||
# Initialize node states
|
||||
for node in workflow.nodes:
|
||||
if node.id not in context.node_states:
|
||||
context.node_states[node.id] = NodeState(node_id=node.id, node_type=node.type, status=NodeStatus.PENDING)
|
||||
|
||||
# Find start node(s)
|
||||
if start_node_id:
|
||||
start_nodes = [node_map[start_node_id]]
|
||||
else:
|
||||
start_nodes = self._find_start_nodes(workflow.nodes, workflow.edges)
|
||||
|
||||
if not start_nodes:
|
||||
raise ValueError('No start nodes found in workflow')
|
||||
|
||||
# Execute from start nodes
|
||||
for start_node in start_nodes:
|
||||
await self._execute_from_node(
|
||||
start_node, node_map, edge_map, context, workflow.settings.max_retries, path=set()
|
||||
)
|
||||
|
||||
# Check final status
|
||||
all_completed = all(
|
||||
state.status in (NodeStatus.COMPLETED, NodeStatus.SKIPPED) for state in context.node_states.values()
|
||||
)
|
||||
|
||||
if all_completed:
|
||||
context.status = ExecutionStatus.COMPLETED
|
||||
else:
|
||||
# Some nodes might still be waiting
|
||||
has_failed = any(state.status == NodeStatus.FAILED for state in context.node_states.values())
|
||||
if has_failed:
|
||||
context.status = ExecutionStatus.FAILED
|
||||
|
||||
except Exception as e:
|
||||
context.status = ExecutionStatus.FAILED
|
||||
context.error = str(e)
|
||||
logger.error(
|
||||
'Workflow execution failed',
|
||||
exc_info=True,
|
||||
extra={
|
||||
'workflow_id': workflow.uuid,
|
||||
'execution_id': context.execution_id,
|
||||
'node_states': {
|
||||
node_id: {
|
||||
'status': state.status.value if state.status else None,
|
||||
'error': state.error,
|
||||
}
|
||||
for node_id, state in context.node_states.items()
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
# Note: Frontend panel logging has been removed.
|
||||
# A new solution will be implemented separately.
|
||||
|
||||
finally:
|
||||
context.end_time = datetime.now()
|
||||
|
||||
# Note: Frontend panel logging has been removed.
|
||||
# A new solution will be implemented separately.
|
||||
|
||||
return context
|
||||
|
||||
async def _execute_from_node(
|
||||
self,
|
||||
node: NodeDefinition,
|
||||
node_map: dict[str, NodeDefinition],
|
||||
edge_map: dict[str, list[EdgeDefinition]],
|
||||
context: ExecutionContext,
|
||||
max_retries: int = 3,
|
||||
path: set[str] | None = None,
|
||||
):
|
||||
"""Execute workflow starting from a specific node"""
|
||||
|
||||
# Initialize path set for cycle detection (path-based, not global visited)
|
||||
if path is None:
|
||||
path = set()
|
||||
|
||||
# Check for circular dependency on the *current path* only
|
||||
# This correctly allows diamond shapes (A→B, A→C, B→D, C→D)
|
||||
if node.id in path:
|
||||
logger.warning(f'Circular dependency detected at node: {node.id}')
|
||||
context.node_states[node.id].status = NodeStatus.SKIPPED
|
||||
context.node_states[node.id].error = 'Circular dependency detected'
|
||||
context.node_states[node.id].end_time = datetime.now()
|
||||
await self._persist_node_execution(node, context.node_states[node.id], context)
|
||||
return
|
||||
|
||||
# Add node to current path
|
||||
path.add(node.id)
|
||||
|
||||
# Check if node should be skipped
|
||||
if await self._should_skip_node(node, context):
|
||||
existing_state = context.node_states[node.id]
|
||||
if existing_state.status == NodeStatus.SKIPPED:
|
||||
existing_state.end_time = existing_state.end_time or datetime.now()
|
||||
await self._persist_node_execution(node, existing_state, context)
|
||||
path.discard(node.id)
|
||||
return
|
||||
|
||||
# Execute current node
|
||||
await self._execute_node(node, context, max_retries)
|
||||
|
||||
# If node failed and we should stop on error, return
|
||||
if context.node_states[node.id].status == NodeStatus.FAILED:
|
||||
path.discard(node.id)
|
||||
return
|
||||
|
||||
node_state = context.node_states[node.id]
|
||||
node_type_name = node.type.split('.')[-1] if '.' in node.type else node.type
|
||||
|
||||
# ── Control flow integration ────────────────────────────────
|
||||
# For loop / iterator nodes: run the LoopExecutor over
|
||||
# downstream body nodes for each item, then continue to the
|
||||
# "completed" output edge.
|
||||
if node_type_name in ('loop', 'iterator'):
|
||||
items = node_state.outputs.get('_items') or []
|
||||
if not items:
|
||||
# iterator: items come from inputs
|
||||
items = node_state.inputs.get('items', node_state.inputs.get('array', []))
|
||||
if not isinstance(items, list):
|
||||
items = [items] if items else []
|
||||
max_iter = int(node.config.get('max_iterations', 100))
|
||||
items = items[:max_iter]
|
||||
|
||||
# Collect downstream "body" nodes (connected via edges)
|
||||
outgoing_edges = edge_map.get(node.id, [])
|
||||
body_nodes = []
|
||||
for edge in outgoing_edges:
|
||||
target = node_map.get(edge.target_node)
|
||||
if target:
|
||||
body_nodes.append(target)
|
||||
|
||||
if body_nodes and items:
|
||||
loop_exec = LoopExecutor(self)
|
||||
results = await loop_exec.execute_loop(items, body_nodes, context, max_iter)
|
||||
node_state.outputs['results'] = results
|
||||
node_state.outputs['completed'] = True
|
||||
else:
|
||||
node_state.outputs['results'] = []
|
||||
node_state.outputs['completed'] = True
|
||||
|
||||
path.discard(node.id)
|
||||
return # body nodes already executed by LoopExecutor
|
||||
|
||||
# For parallel nodes: run downstream branches concurrently
|
||||
if node_type_name == 'parallel':
|
||||
outgoing_edges = edge_map.get(node.id, [])
|
||||
branch_nodes = []
|
||||
for edge in outgoing_edges:
|
||||
target = node_map.get(edge.target_node)
|
||||
if target:
|
||||
branch_nodes.append([target])
|
||||
|
||||
if branch_nodes:
|
||||
par_exec = ParallelExecutor(self)
|
||||
results = await par_exec.execute_parallel(branch_nodes, context)
|
||||
node_state.outputs['results'] = results
|
||||
|
||||
path.discard(node.id)
|
||||
return # branch nodes already executed by ParallelExecutor
|
||||
|
||||
# ── Standard edge-based continuation ────────────────────────
|
||||
# Get outgoing edges
|
||||
outgoing_edges = edge_map.get(node.id, [])
|
||||
|
||||
# Execute next nodes based on edge conditions
|
||||
for edge in outgoing_edges:
|
||||
target_node = node_map.get(edge.target_node)
|
||||
if not target_node:
|
||||
continue
|
||||
|
||||
# Check edge condition
|
||||
if edge.condition:
|
||||
condition_met = await self._evaluate_condition(edge.condition, context)
|
||||
if not condition_met:
|
||||
continue
|
||||
|
||||
# Check if all inputs are ready
|
||||
if await self._inputs_ready(target_node, edge_map, context):
|
||||
await self._execute_from_node(target_node, node_map, edge_map, context, max_retries, path)
|
||||
|
||||
# Remove node from path when backtracking (allows diamond revisit)
|
||||
path.discard(node.id)
|
||||
|
||||
async def _execute_node(self, node: NodeDefinition, context: ExecutionContext, max_retries: int = 3):
|
||||
"""Execute a single node with retry logic"""
|
||||
|
||||
node_state = context.node_states[node.id]
|
||||
node_state.status = NodeStatus.RUNNING
|
||||
node_state.start_time = datetime.now()
|
||||
|
||||
# Get node instance (pass ap for access to services)
|
||||
node_instance = self.registry.create_instance(node.type, node.id, node.config, ap=self.ap)
|
||||
|
||||
if not node_instance:
|
||||
node_state.status = NodeStatus.FAILED
|
||||
node_state.error = f'Unknown node type: {node.type}'
|
||||
node_state.end_time = datetime.now()
|
||||
self._record_execution_step(node, node_state, context)
|
||||
await self._persist_node_execution(node, node_state, context)
|
||||
return
|
||||
|
||||
# Resolve inputs
|
||||
inputs = await self._resolve_inputs(node, context)
|
||||
node_state.inputs = inputs
|
||||
|
||||
# Validate inputs
|
||||
validation_errors = await node_instance.validate_inputs(inputs)
|
||||
if validation_errors:
|
||||
node_state.status = NodeStatus.FAILED
|
||||
node_state.error = '; '.join(validation_errors)
|
||||
node_state.end_time = datetime.now()
|
||||
self._record_execution_step(node, node_state, context)
|
||||
await self._persist_node_execution(node, node_state, context)
|
||||
return
|
||||
|
||||
# Check if node supports streaming (has execute_stream method and stream config is enabled)
|
||||
use_streaming = hasattr(node_instance, 'execute_stream') and node.config.get('stream', False)
|
||||
|
||||
# Execute with retries
|
||||
for attempt in range(max_retries + 1):
|
||||
try:
|
||||
if use_streaming:
|
||||
# Streaming execution with aggregation and timeout
|
||||
aggregated_response = ''
|
||||
try:
|
||||
async with asyncio.timeout(300): # 5 minute timeout for streaming
|
||||
async for chunk in node_instance.execute_stream(inputs, context):
|
||||
if chunk:
|
||||
aggregated_response += chunk
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(f'Node {node.id} ({node.type}) streaming timed out, falling back to non-streaming')
|
||||
use_streaming = False
|
||||
outputs = await node_instance.execute(inputs, context)
|
||||
else:
|
||||
# Get response from context if set by execute_stream, otherwise use aggregated
|
||||
final_response = context.variables.pop('_last_llm_response', aggregated_response)
|
||||
outputs = {'response': final_response, 'usage': {'prompt_tokens': 0, 'completion_tokens': 0, 'total_tokens': 0}}
|
||||
logger.info(f'Node {node.id} ({node.type}) streaming completed, response length: {len(final_response)}')
|
||||
else:
|
||||
outputs = await node_instance.execute(inputs, context)
|
||||
node_state.outputs = outputs
|
||||
node_state.status = NodeStatus.COMPLETED
|
||||
node_state.end_time = datetime.now()
|
||||
break
|
||||
except Exception as e:
|
||||
node_state.retry_count = attempt + 1
|
||||
logger.error(
|
||||
f'Node {node.id} ({node.type}) execution failed (attempt {attempt + 1}/{max_retries + 1}): {e}',
|
||||
exc_info=True,
|
||||
extra={
|
||||
'node_id': node.id,
|
||||
'node_type': node.type,
|
||||
'attempt': attempt + 1,
|
||||
'max_retries': max_retries,
|
||||
'execution_id': context.execution_id,
|
||||
},
|
||||
)
|
||||
|
||||
if attempt < max_retries:
|
||||
await asyncio.sleep(1) # Brief delay before retry
|
||||
else:
|
||||
node_state.status = NodeStatus.FAILED
|
||||
node_state.error = str(e)
|
||||
node_state.end_time = datetime.now()
|
||||
logger.error(
|
||||
f'Node {node.id} ({node.type}) permanently failed after {max_retries + 1} attempts',
|
||||
extra={
|
||||
'node_id': node.id,
|
||||
'node_type': node.type,
|
||||
'error': str(e),
|
||||
'execution_id': context.execution_id,
|
||||
},
|
||||
)
|
||||
|
||||
self._record_execution_step(node, node_state, context)
|
||||
await self._persist_node_execution(node, node_state, context)
|
||||
|
||||
async def _resolve_inputs(self, node: NodeDefinition, context: ExecutionContext) -> dict[str, Any]:
|
||||
"""Resolve input values for a node from connected nodes and context"""
|
||||
inputs = {}
|
||||
|
||||
# Get inputs from context variables
|
||||
if 'message' in context.variables:
|
||||
inputs['message'] = context.variables['message']
|
||||
|
||||
# Get inputs from message context
|
||||
if context.message_context:
|
||||
inputs['message'] = context.message_context.message_content
|
||||
inputs['message_content'] = context.message_context.message_content
|
||||
inputs['sender_id'] = context.message_context.sender_id
|
||||
inputs['platform'] = context.message_context.platform
|
||||
else:
|
||||
logger.warning(
|
||||
f'[_resolve_inputs] node={node.id} ({node.type}): message_context is None!',
|
||||
extra={
|
||||
'node_id': node.id,
|
||||
'node_type': node.type,
|
||||
'execution_id': context.execution_id,
|
||||
'variables_keys': list(context.variables.keys()) if context.variables else [],
|
||||
},
|
||||
)
|
||||
|
||||
# Log current inputs state after message_context processing
|
||||
logger.debug(
|
||||
f'[_resolve_inputs] node={node.id} after message_context: {list(inputs.keys())}',
|
||||
)
|
||||
|
||||
# Get inputs from node config that reference other nodes
|
||||
for key, value in node.config.items():
|
||||
if isinstance(value, str) and value.startswith('{{') and value.endswith('}}'):
|
||||
resolved = await self._resolve_expression(value[2:-2], context)
|
||||
inputs[key] = resolved
|
||||
else:
|
||||
inputs[key] = value
|
||||
|
||||
# Get inputs from connected upstream nodes via data edges.
|
||||
# Build a reverse map: for each incoming edge to this node, find the
|
||||
# source node and the specific source/target port.
|
||||
for edge in self._edges:
|
||||
if not self._is_data_edge(edge):
|
||||
continue
|
||||
if edge.target_node != node.id:
|
||||
continue
|
||||
source_state = context.node_states.get(edge.source_node)
|
||||
if not source_state or source_state.status != NodeStatus.COMPLETED:
|
||||
continue
|
||||
target_port = edge.target_port or 'input'
|
||||
source_port = edge.source_port or 'output'
|
||||
# Map the source node's output port value to this node's input port
|
||||
if source_port in source_state.outputs:
|
||||
inputs[target_port] = source_state.outputs[source_port]
|
||||
elif 'output' in source_state.outputs:
|
||||
# Fallback: if exact port not found, try generic 'output'
|
||||
inputs[target_port] = source_state.outputs['output']
|
||||
elif source_state.outputs:
|
||||
# Last resort: use the first available output
|
||||
inputs[target_port] = next(iter(source_state.outputs.values()))
|
||||
|
||||
# Smart input mapping: if a node needs 'message' but received a different
|
||||
# port name (e.g., 'content' from llm_call), copy the value to 'message'.
|
||||
# This handles edge connection mismatches where the sender uses a different
|
||||
# port name than what the receiver expects.
|
||||
if 'message' not in inputs or inputs.get('message') is None:
|
||||
for fallback_key in ('content', 'response', 'input', 'output', 'result', 'text'):
|
||||
if fallback_key in inputs and inputs[fallback_key] is not None:
|
||||
inputs['message'] = inputs[fallback_key]
|
||||
logger.debug(
|
||||
f'[_resolve_inputs] node={node.id}: mapped {fallback_key} -> message',
|
||||
)
|
||||
break
|
||||
|
||||
logger.debug(
|
||||
f'[_resolve_inputs] node={node.id} final inputs keys: {list(inputs.keys())}, message={repr(inputs.get("message", "<missing>")[:100] if isinstance(inputs.get("message"), str) else inputs.get("message"))}',
|
||||
)
|
||||
return inputs
|
||||
|
||||
async def _resolve_expression(self, expression: str, context: ExecutionContext) -> Any:
|
||||
"""Resolve a variable expression like 'nodes.node1.outputs.text'"""
|
||||
parts = expression.strip().split('.')
|
||||
|
||||
if not parts:
|
||||
return None
|
||||
|
||||
if parts[0] == 'nodes' and len(parts) >= 4:
|
||||
# nodes.node_id.outputs.output_name
|
||||
node_id = parts[1]
|
||||
if parts[2] == 'outputs' and node_id in context.node_states:
|
||||
output_name = '.'.join(parts[3:])
|
||||
return context.node_states[node_id].outputs.get(output_name)
|
||||
|
||||
elif parts[0] == 'variables':
|
||||
# variables.var_name
|
||||
var_name = '.'.join(parts[1:])
|
||||
return context.variables.get(var_name)
|
||||
|
||||
elif parts[0] == 'conversation_variables':
|
||||
# conversation_variables.var_name
|
||||
var_name = '.'.join(parts[1:])
|
||||
return context.conversation_variables.get(var_name)
|
||||
|
||||
elif parts[0] == 'message':
|
||||
# message.content, message.sender_id, etc.
|
||||
if context.message_context:
|
||||
attr = parts[1] if len(parts) > 1 else None
|
||||
if attr == 'content':
|
||||
return context.message_context.message_content
|
||||
elif attr == 'sender_id':
|
||||
return context.message_context.sender_id
|
||||
elif attr == 'platform':
|
||||
return context.message_context.platform
|
||||
elif attr == 'conversation_id':
|
||||
return context.message_context.conversation_id
|
||||
|
||||
return None
|
||||
|
||||
async def _evaluate_condition(self, condition: str, context: ExecutionContext) -> bool:
|
||||
"""Evaluate a condition expression safely.
|
||||
|
||||
Any ``{{ ... }}`` references are resolved against the execution context
|
||||
and bound as **variables** that are passed to :func:`safe_eval_with_vars`.
|
||||
Values are never string-concatenated into the expression, which avoids
|
||||
broken parsing (e.g. values containing quotes) and any injection risk
|
||||
from non-literal value types (lists, dicts, etc.).
|
||||
"""
|
||||
variables: dict[str, Any] = {}
|
||||
try:
|
||||
# Resolve variable references in condition into bound variables.
|
||||
if '{{' in condition:
|
||||
pattern = r'\{\{([^}]+)\}\}'
|
||||
|
||||
placeholders: dict[str, str] = {}
|
||||
placeholder_idx = 0
|
||||
|
||||
def replace_with_placeholder(match: re.Match[str]) -> str:
|
||||
nonlocal placeholder_idx
|
||||
var_expr = match.group(1)
|
||||
placeholder = f'__ph{placeholder_idx}__'
|
||||
placeholders[placeholder] = var_expr
|
||||
placeholder_idx += 1
|
||||
return placeholder
|
||||
|
||||
condition = re.sub(pattern, replace_with_placeholder, condition)
|
||||
|
||||
# Resolve each placeholder and bind it as a variable, so the
|
||||
# actual value (of any type) is passed through unchanged.
|
||||
for placeholder, var_expr in placeholders.items():
|
||||
variables[placeholder] = await self._resolve_expression(var_expr, context)
|
||||
|
||||
# Safe expression evaluation with bound variables (AST whitelist).
|
||||
result = safe_eval_with_vars(condition, variables)
|
||||
return bool(result)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f'Condition evaluation failed: {condition} - {e}')
|
||||
return False
|
||||
|
||||
async def _should_skip_node(self, node: NodeDefinition, context: ExecutionContext) -> bool:
|
||||
"""Check if a node should be skipped"""
|
||||
state = context.node_states.get(node.id)
|
||||
if state and state.status in (NodeStatus.COMPLETED, NodeStatus.RUNNING, NodeStatus.SKIPPED):
|
||||
return True
|
||||
return False
|
||||
|
||||
async def _inputs_ready(
|
||||
self, node: NodeDefinition, edge_map: dict[str, list[EdgeDefinition]], context: ExecutionContext
|
||||
) -> bool:
|
||||
"""Check if all control predecessors and data providers are ready."""
|
||||
incoming_nodes = self._incoming_dependency_nodes(node.id, edge_map)
|
||||
|
||||
# Check if all incoming nodes have completed
|
||||
for source_id in incoming_nodes:
|
||||
state = context.node_states.get(source_id)
|
||||
if not state or state.status not in (NodeStatus.COMPLETED, NodeStatus.SKIPPED):
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def _find_start_nodes(self, nodes: list[NodeDefinition], edges: list[EdgeDefinition]) -> list[NodeDefinition]:
|
||||
"""Find nodes that have no incoming edges (start nodes)"""
|
||||
target_nodes = {edge.target_node for edge in edges if self._is_control_edge(edge) or self._is_data_edge(edge)}
|
||||
start_nodes = [node for node in nodes if node.id not in target_nodes]
|
||||
|
||||
# Also check for trigger nodes
|
||||
trigger_types = {'message_trigger', 'cron_trigger', 'webhook_trigger', 'event_trigger'}
|
||||
for node in nodes:
|
||||
if node.type in trigger_types and node not in start_nodes:
|
||||
start_nodes.insert(0, node)
|
||||
|
||||
return start_nodes
|
||||
|
||||
def _build_edge_map(self, edges: list[EdgeDefinition]) -> dict[str, list[EdgeDefinition]]:
|
||||
"""Build a map of source node ID to outgoing control edges."""
|
||||
edge_map: dict[str, list[EdgeDefinition]] = {}
|
||||
for edge in edges:
|
||||
if not self._is_control_edge(edge):
|
||||
continue
|
||||
if edge.source_node not in edge_map:
|
||||
edge_map[edge.source_node] = []
|
||||
edge_map[edge.source_node].append(edge)
|
||||
return edge_map
|
||||
|
||||
def _edge_type(self, edge: EdgeDefinition) -> str:
|
||||
edge_type = (getattr(edge, 'edge_type', None) or 'legacy').strip().lower()
|
||||
if edge_type not in {'control', 'data', 'legacy'}:
|
||||
return 'legacy'
|
||||
return edge_type
|
||||
|
||||
def _is_control_edge(self, edge: EdgeDefinition) -> bool:
|
||||
return self._edge_type(edge) in {'control', 'legacy'}
|
||||
|
||||
def _is_data_edge(self, edge: EdgeDefinition) -> bool:
|
||||
return self._edge_type(edge) in {'data', 'legacy'}
|
||||
|
||||
def _incoming_dependency_nodes(
|
||||
self, node_id: str, edge_map: dict[str, list[EdgeDefinition]]
|
||||
) -> set[str]:
|
||||
incoming_nodes: set[str] = set()
|
||||
for source_id, edges in edge_map.items():
|
||||
for edge in edges:
|
||||
if edge.target_node == node_id:
|
||||
incoming_nodes.add(source_id)
|
||||
|
||||
for edge in self._edges:
|
||||
if self._is_data_edge(edge) and edge.target_node == node_id:
|
||||
incoming_nodes.add(edge.source_node)
|
||||
|
||||
return incoming_nodes
|
||||
|
||||
def _record_execution_step(self, node: NodeDefinition, node_state: NodeState, context: ExecutionContext):
|
||||
"""Record an execution step in the history"""
|
||||
duration_ms = 0
|
||||
if node_state.start_time and node_state.end_time:
|
||||
duration_ms = int((node_state.end_time - node_state.start_time).total_seconds() * 1000)
|
||||
|
||||
step = ExecutionStep(
|
||||
step_id=f"step_{uuid.uuid4().hex[:8]}",
|
||||
timestamp=datetime.now(),
|
||||
node_id=node.id,
|
||||
node_type=node.type,
|
||||
status=node_state.status,
|
||||
duration_ms=duration_ms,
|
||||
error=node_state.error,
|
||||
inputs=node_state.inputs,
|
||||
outputs=node_state.outputs,
|
||||
)
|
||||
context.history.append(step)
|
||||
|
||||
async def _persist_node_execution(
|
||||
self,
|
||||
node: NodeDefinition,
|
||||
node_state: NodeState,
|
||||
context: ExecutionContext,
|
||||
):
|
||||
"""Persist node execution state for execution detail and logs."""
|
||||
if not self.ap:
|
||||
return
|
||||
|
||||
values = {
|
||||
'execution_uuid': context.execution_id,
|
||||
'node_id': node.id,
|
||||
'node_type': node.type,
|
||||
'status': node_state.status.value,
|
||||
'inputs': node_state.inputs,
|
||||
'outputs': node_state.outputs,
|
||||
'start_time': node_state.start_time,
|
||||
'end_time': node_state.end_time,
|
||||
'error': node_state.error,
|
||||
'retry_count': node_state.retry_count,
|
||||
}
|
||||
|
||||
existing_query = sqlalchemy.select(persistence_workflow.WorkflowNodeExecution).where(
|
||||
persistence_workflow.WorkflowNodeExecution.execution_uuid == context.execution_id,
|
||||
persistence_workflow.WorkflowNodeExecution.node_id == node.id,
|
||||
)
|
||||
existing_result = await self.ap.persistence_mgr.execute_async(existing_query)
|
||||
existing = existing_result.first()
|
||||
|
||||
if existing is None:
|
||||
await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.insert(persistence_workflow.WorkflowNodeExecution).values(**values)
|
||||
)
|
||||
else:
|
||||
await self.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.update(persistence_workflow.WorkflowNodeExecution)
|
||||
.where(persistence_workflow.WorkflowNodeExecution.id == existing.id)
|
||||
.values(**values)
|
||||
)
|
||||
|
||||
|
||||
class ParallelExecutor:
|
||||
"""Execute multiple branches in parallel"""
|
||||
|
||||
def __init__(self, executor: WorkflowExecutor):
|
||||
self.executor = executor
|
||||
|
||||
async def execute_parallel(
|
||||
self, branches: list[list[NodeDefinition]], context: ExecutionContext
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Execute multiple branches in parallel.
|
||||
|
||||
Args:
|
||||
branches: List of node sequences to execute in parallel
|
||||
context: Execution context
|
||||
|
||||
Returns:
|
||||
List of results from each branch
|
||||
"""
|
||||
tasks = []
|
||||
for branch in branches:
|
||||
task = self._execute_branch(branch, context)
|
||||
tasks.append(task)
|
||||
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
processed_results = []
|
||||
for index, result in enumerate(results):
|
||||
if isinstance(result, Exception):
|
||||
logger.error(
|
||||
f'Parallel branch {index} failed: {result}',
|
||||
exc_info=result,
|
||||
extra={'branch_index': index, 'execution_id': context.execution_id},
|
||||
)
|
||||
processed_results.append({'error': str(result)})
|
||||
else:
|
||||
processed_results.append(result)
|
||||
|
||||
return processed_results
|
||||
|
||||
async def _execute_branch(self, nodes: list[NodeDefinition], context: ExecutionContext) -> dict[str, Any]:
|
||||
"""Execute a single branch"""
|
||||
# Create a copy of context for this branch
|
||||
branch_outputs = {}
|
||||
|
||||
for node in nodes:
|
||||
await self.executor._execute_node(node, context, max_retries=3)
|
||||
state = context.node_states.get(node.id)
|
||||
if state and state.status == NodeStatus.COMPLETED:
|
||||
branch_outputs[node.id] = state.outputs
|
||||
elif state and state.status == NodeStatus.FAILED:
|
||||
branch_outputs['error'] = state.error
|
||||
break
|
||||
|
||||
return branch_outputs
|
||||
|
||||
|
||||
class LoopExecutor:
|
||||
"""Execute loop iterations"""
|
||||
|
||||
def __init__(self, executor: WorkflowExecutor):
|
||||
self.executor = executor
|
||||
|
||||
async def execute_loop(
|
||||
self, items: list[Any], loop_body: list[NodeDefinition], context: ExecutionContext, max_iterations: int = 100
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Execute a loop over items.
|
||||
|
||||
Args:
|
||||
items: Items to iterate over
|
||||
loop_body: Nodes to execute for each item
|
||||
context: Execution context
|
||||
max_iterations: Maximum number of iterations
|
||||
|
||||
Returns:
|
||||
List of results from each iteration
|
||||
"""
|
||||
results = []
|
||||
|
||||
for i, item in enumerate(items[:max_iterations]):
|
||||
# Set loop variables
|
||||
context.variables['loop_item'] = item
|
||||
context.variables['loop_index'] = i
|
||||
context.variables['loop_is_first'] = i == 0
|
||||
context.variables['loop_is_last'] = i == len(items) - 1
|
||||
|
||||
iteration_result = {}
|
||||
|
||||
for node in loop_body:
|
||||
# Reset node state for this iteration
|
||||
context.node_states[node.id] = NodeState(node_id=node.id, node_type=node.type, status=NodeStatus.PENDING)
|
||||
|
||||
await self.executor._execute_node(node, context, max_retries=3)
|
||||
|
||||
state = context.node_states.get(node.id)
|
||||
if state:
|
||||
iteration_result[node.id] = state.outputs
|
||||
|
||||
# Check for break condition
|
||||
if state.outputs.get('break', False):
|
||||
results.append(iteration_result)
|
||||
return results
|
||||
|
||||
results.append(iteration_result)
|
||||
|
||||
# Clean up loop variables
|
||||
context.variables.pop('loop_item', None)
|
||||
context.variables.pop('loop_index', None)
|
||||
context.variables.pop('loop_is_first', None)
|
||||
context.variables.pop('loop_is_last', None)
|
||||
|
||||
return results
|
||||
@@ -0,0 +1,284 @@
|
||||
"""Workflow node metadata loading and validation.
|
||||
|
||||
This module makes YAML files under ``templates/metadata/nodes`` the backend
|
||||
source of truth for workflow node metadata. Python node classes still provide
|
||||
execution logic, but UI-facing metadata is loaded from YAML.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import logging
|
||||
from importlib import resources
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional
|
||||
|
||||
import yaml
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MetadataLoadError(Exception):
|
||||
"""Raised when a workflow node metadata file cannot be loaded."""
|
||||
|
||||
|
||||
class MetadataValidationError(Exception):
|
||||
"""Raised when workflow node metadata does not match the expected shape."""
|
||||
|
||||
|
||||
class NodeMetadataValidator:
|
||||
"""Validate workflow node metadata loaded from YAML files.
|
||||
|
||||
The validator is intentionally strict about the structural fields that the
|
||||
editor needs, but tolerant of legacy YAML details such as missing top-level
|
||||
``label`` or additional frontend field types.
|
||||
"""
|
||||
|
||||
REQUIRED_FIELDS = ('name', 'category', 'inputs', 'outputs', 'config')
|
||||
VALID_CATEGORIES = {'trigger', 'process', 'control', 'action', 'integration', 'misc'}
|
||||
VALID_PORT_TYPES = {'any', 'string', 'number', 'integer', 'boolean', 'object', 'array', 'datetime', 'null'}
|
||||
VALID_CONFIG_TYPES = {
|
||||
'string',
|
||||
'integer',
|
||||
'number',
|
||||
'float',
|
||||
'boolean',
|
||||
'select',
|
||||
'json',
|
||||
'textarea',
|
||||
'text',
|
||||
'secret',
|
||||
'array[string]',
|
||||
'file',
|
||||
'array[file]',
|
||||
'llm-model-selector',
|
||||
'embedding-model-selector',
|
||||
'rerank-model-selector',
|
||||
'pipeline-selector',
|
||||
'knowledge-base-selector',
|
||||
'knowledge-base-multi-selector',
|
||||
'bot-selector',
|
||||
'tools-selector',
|
||||
'model-fallback-selector',
|
||||
'prompt-editor',
|
||||
'plugin-selector',
|
||||
'webhook-url',
|
||||
'embed-code',
|
||||
'workflow-selector',
|
||||
}
|
||||
|
||||
def validate(self, metadata: dict[str, Any]) -> list[str]:
|
||||
"""Return validation errors. An empty list means the metadata is valid."""
|
||||
errors: list[str] = []
|
||||
|
||||
if not isinstance(metadata, dict):
|
||||
return ['metadata root must be a mapping']
|
||||
|
||||
for field in self.REQUIRED_FIELDS:
|
||||
if field not in metadata:
|
||||
errors.append(f'missing required field: {field}')
|
||||
|
||||
if errors:
|
||||
return errors
|
||||
|
||||
name = metadata.get('name')
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
errors.append('field "name" must be a non-empty string')
|
||||
|
||||
category = metadata.get('category')
|
||||
if category not in self.VALID_CATEGORIES:
|
||||
errors.append(f'invalid category: {category}')
|
||||
|
||||
errors.extend(self._validate_ports(metadata.get('inputs'), 'inputs'))
|
||||
errors.extend(self._validate_ports(metadata.get('outputs'), 'outputs'))
|
||||
errors.extend(self._validate_config(metadata.get('config')))
|
||||
|
||||
return errors
|
||||
|
||||
def validate_or_raise(self, metadata: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Validate metadata and raise ``MetadataValidationError`` on failure."""
|
||||
errors = self.validate(metadata)
|
||||
if errors:
|
||||
node_name = metadata.get('name', 'unknown') if isinstance(metadata, dict) else 'unknown'
|
||||
raise MetadataValidationError(f'invalid metadata for {node_name}: {errors}')
|
||||
return metadata
|
||||
|
||||
def _validate_ports(self, ports: Any, field_name: str) -> list[str]:
|
||||
errors: list[str] = []
|
||||
if not isinstance(ports, list):
|
||||
return [f'{field_name} must be a list']
|
||||
|
||||
seen_names: set[str] = set()
|
||||
for index, port in enumerate(ports):
|
||||
path = f'{field_name}[{index}]'
|
||||
if not isinstance(port, dict):
|
||||
errors.append(f'{path} must be a mapping')
|
||||
continue
|
||||
|
||||
name = port.get('name')
|
||||
if not isinstance(name, str) or not name:
|
||||
errors.append(f'{path}.name must be a non-empty string')
|
||||
continue
|
||||
|
||||
if name in seen_names:
|
||||
errors.append(f'{path}.name duplicates "{name}"')
|
||||
seen_names.add(name)
|
||||
|
||||
port_type = port.get('type', 'any')
|
||||
if port_type not in self.VALID_PORT_TYPES:
|
||||
errors.append(f'{path}.type has unsupported value "{port_type}"')
|
||||
|
||||
return errors
|
||||
|
||||
def _validate_config(self, config: Any) -> list[str]:
|
||||
errors: list[str] = []
|
||||
if not isinstance(config, list):
|
||||
return ['config must be a list']
|
||||
|
||||
seen_names: set[str] = set()
|
||||
for index, item in enumerate(config):
|
||||
path = f'config[{index}]'
|
||||
if not isinstance(item, dict):
|
||||
errors.append(f'{path} must be a mapping')
|
||||
continue
|
||||
|
||||
name = item.get('name')
|
||||
if not isinstance(name, str) or not name:
|
||||
errors.append(f'{path}.name must be a non-empty string')
|
||||
continue
|
||||
|
||||
if name in seen_names:
|
||||
errors.append(f'{path}.name duplicates "{name}"')
|
||||
seen_names.add(name)
|
||||
|
||||
item_type = item.get('type', 'string')
|
||||
if item_type not in self.VALID_CONFIG_TYPES:
|
||||
errors.append(f'{path}.type has unsupported value "{item_type}"')
|
||||
|
||||
min_value = item.get('min_value')
|
||||
max_value = item.get('max_value')
|
||||
if isinstance(min_value, (int, float)) and isinstance(max_value, (int, float)) and min_value > max_value:
|
||||
errors.append(f'{path}.min_value must be <= max_value')
|
||||
|
||||
return errors
|
||||
|
||||
|
||||
class NodeMetadataLoader:
|
||||
"""Load and cache workflow node metadata from YAML files."""
|
||||
|
||||
def __init__(self, validator: Optional[NodeMetadataValidator] = None) -> None:
|
||||
self._validator = validator or NodeMetadataValidator()
|
||||
self._metadata: dict[str, dict[str, Any]] = {}
|
||||
self._sources: dict[str, str] = {}
|
||||
self._load_errors: list[dict[str, str]] = []
|
||||
|
||||
async def load_core_metadata(self, resource_dir: str = 'metadata/nodes') -> int:
|
||||
"""Load all core node metadata from the ``langbot.templates`` package."""
|
||||
return await self.load_package_directory('langbot.templates', resource_dir, source='core')
|
||||
|
||||
async def load_package_directory(self, package: str, resource_dir: str, source: str = 'core') -> int:
|
||||
"""Load YAML files from a package resource directory."""
|
||||
try:
|
||||
root = resources.files(package).joinpath(resource_dir)
|
||||
yaml_files = sorted(
|
||||
(item for item in root.iterdir() if item.is_file() and item.name.endswith(('.yaml', '.yml'))),
|
||||
key=lambda item: item.name,
|
||||
)
|
||||
except Exception as exc:
|
||||
raise MetadataLoadError(f'failed to scan package directory {package}:{resource_dir}: {exc}') from exc
|
||||
|
||||
return self._load_files(yaml_files, source=source)
|
||||
|
||||
async def load_directory(self, directory: str | Path, source: str) -> int:
|
||||
"""Load YAML files from an external filesystem directory, e.g. a plugin."""
|
||||
directory_path = Path(directory)
|
||||
if not directory_path.exists():
|
||||
logger.warning('Workflow metadata directory does not exist: %s', directory_path)
|
||||
return 0
|
||||
if not directory_path.is_dir():
|
||||
raise MetadataLoadError(f'workflow metadata path is not a directory: {directory_path}')
|
||||
|
||||
yaml_files = sorted(directory_path.glob('*.yml')) + sorted(directory_path.glob('*.yaml'))
|
||||
return self._load_files(yaml_files, source=source)
|
||||
|
||||
def get_metadata(self, node_type: str) -> Optional[dict[str, Any]]:
|
||||
"""Return metadata by full type or short node name."""
|
||||
if node_type in self._metadata:
|
||||
return copy.deepcopy(self._metadata[node_type])
|
||||
|
||||
short_name = node_type.split('.')[-1]
|
||||
for registered_type, metadata in self._metadata.items():
|
||||
if registered_type.split('.')[-1] == short_name or metadata.get('name') == short_name:
|
||||
return copy.deepcopy(metadata)
|
||||
|
||||
return None
|
||||
|
||||
def get_all_metadata(self) -> dict[str, dict[str, Any]]:
|
||||
"""Return a deep copy of all loaded metadata keyed by canonical node type."""
|
||||
return copy.deepcopy(self._metadata)
|
||||
|
||||
def get_load_errors(self) -> list[dict[str, str]]:
|
||||
"""Return metadata files that failed to load or validate."""
|
||||
return copy.deepcopy(self._load_errors)
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Clear all cached metadata and errors."""
|
||||
self._metadata.clear()
|
||||
self._sources.clear()
|
||||
self._load_errors.clear()
|
||||
|
||||
def _load_files(self, yaml_files: Iterable[Any], source: str) -> int:
|
||||
count = 0
|
||||
for yaml_file in yaml_files:
|
||||
file_name = getattr(yaml_file, 'name', str(yaml_file))
|
||||
try:
|
||||
metadata = self._load_yaml(yaml_file)
|
||||
self._validator.validate_or_raise(metadata)
|
||||
node_type = build_node_type(metadata)
|
||||
|
||||
if node_type in self._metadata:
|
||||
existing_source = self._sources.get(node_type, 'unknown')
|
||||
if existing_source == 'core' and source != 'core':
|
||||
raise MetadataLoadError(
|
||||
f'plugin source "{source}" attempted to override core node "{node_type}"'
|
||||
)
|
||||
logger.warning(
|
||||
'Workflow node metadata %s from %s overrides previous source %s',
|
||||
node_type,
|
||||
source,
|
||||
existing_source,
|
||||
)
|
||||
|
||||
cached_metadata = copy.deepcopy(metadata)
|
||||
cached_metadata['_source'] = source
|
||||
cached_metadata['_file'] = file_name
|
||||
self._metadata[node_type] = cached_metadata
|
||||
self._sources[node_type] = source
|
||||
count += 1
|
||||
except Exception as exc:
|
||||
self._load_errors.append({'file': file_name, 'source': source, 'error': str(exc)})
|
||||
logger.error('Failed to load workflow node metadata %s: %s', file_name, exc)
|
||||
|
||||
return count
|
||||
|
||||
def _load_yaml(self, yaml_file: Any) -> dict[str, Any]:
|
||||
try:
|
||||
if hasattr(yaml_file, 'open'):
|
||||
with yaml_file.open('r', encoding='utf-8') as file:
|
||||
data = yaml.load(file, Loader=yaml.FullLoader)
|
||||
else:
|
||||
with open(yaml_file, 'r', encoding='utf-8') as file:
|
||||
data = yaml.load(file, Loader=yaml.FullLoader)
|
||||
except Exception as exc:
|
||||
raise MetadataLoadError(f'failed to parse YAML: {exc}') from exc
|
||||
|
||||
if not isinstance(data, dict):
|
||||
raise MetadataLoadError('YAML root must be a mapping')
|
||||
return data
|
||||
|
||||
|
||||
def build_node_type(metadata: dict[str, Any]) -> str:
|
||||
"""Build canonical ``category.name`` node type from metadata."""
|
||||
category = metadata.get('category') or 'misc'
|
||||
name = metadata.get('name') or ''
|
||||
return f'{category}.{name}'
|
||||
@@ -0,0 +1,61 @@
|
||||
"""
|
||||
Monitoring helper for recording events during workflow execution.
|
||||
This module provides convenient methods to record monitoring data
|
||||
without cluttering the main workflow code.
|
||||
|
||||
NOTE: All frontend panel logging functionality has been removed.
|
||||
A new solution will be implemented separately.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import typing
|
||||
import time
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from ..core import app
|
||||
from langbot_plugin.api.entities.builtin.workflow.query import WorkflowQuery
|
||||
|
||||
|
||||
class WorkflowMonitoringHelper:
|
||||
"""Helper class for workflow monitoring operations"""
|
||||
|
||||
# All frontend panel logging methods have been removed.
|
||||
# A new solution will be implemented separately.
|
||||
pass
|
||||
|
||||
|
||||
class LLMCallMonitor:
|
||||
"""Context manager for monitoring LLM calls in workflow"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ap: app.Application,
|
||||
query: WorkflowQuery,
|
||||
bot_id: str,
|
||||
bot_name: str,
|
||||
workflow_id: str,
|
||||
workflow_name: str,
|
||||
node_name: str,
|
||||
model_name: str,
|
||||
):
|
||||
self.ap = ap
|
||||
self.query = query
|
||||
self.bot_id = bot_id
|
||||
self.bot_name = bot_name
|
||||
self.workflow_id = workflow_id
|
||||
self.workflow_name = workflow_name
|
||||
self.node_name = node_name
|
||||
self.model_name = model_name
|
||||
self.start_time = None
|
||||
self.input_tokens = 0
|
||||
self.output_tokens = 0
|
||||
|
||||
async def __aenter__(self):
|
||||
self.start_time = time.time()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
# LLM call monitoring has been removed.
|
||||
# A new solution will be implemented separately.
|
||||
return False
|
||||
@@ -0,0 +1,363 @@
|
||||
"""
|
||||
Monitoring helper for recording events during workflow execution.
|
||||
This module provides convenient methods to record monitoring data
|
||||
without cluttering the main workflow code.
|
||||
|
||||
Logging scheme (aligned with pipeline monitoring):
|
||||
- Trigger log: stores original user message content directly
|
||||
- LLM call log: uses record_llm_call only (no additional message record)
|
||||
- LLM response log: stores response message content directly
|
||||
- Reply log: stores reply content directly
|
||||
|
||||
Fields are extracted from WorkflowQuery object when available, with fallback to context_vars.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import typing
|
||||
import time
|
||||
import json
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from ..core import app
|
||||
|
||||
|
||||
class WorkflowMonitoringHelper:
|
||||
"""Helper class for workflow monitoring operations"""
|
||||
|
||||
@staticmethod
|
||||
def _is_workflow_query(query) -> bool:
|
||||
"""Check if query is a WorkflowQuery object"""
|
||||
if query is None or isinstance(query, str):
|
||||
return False
|
||||
# Check for WorkflowQuery attributes
|
||||
return hasattr(query, 'launcher_type') or hasattr(query, 'workflow_uuid')
|
||||
|
||||
@staticmethod
|
||||
def _get_session_id(query, context_vars: dict | None = None) -> str:
|
||||
"""Build session_id from query or context_vars"""
|
||||
# Try to get from WorkflowQuery first
|
||||
if WorkflowMonitoringHelper._is_workflow_query(query) and query.launcher_type:
|
||||
launcher_type = query.launcher_type.value if hasattr(query.launcher_type, 'value') else str(query.launcher_type)
|
||||
launcher_id = query.launcher_id or 'unknown'
|
||||
return f'{launcher_type}_{launcher_id}'
|
||||
|
||||
# Fallback to context_vars
|
||||
if context_vars and context_vars.get('_launcher_type') and context_vars.get('_launcher_id'):
|
||||
return f"{context_vars['_launcher_type']}_{context_vars['_launcher_id']}"
|
||||
|
||||
return 'workflow_session'
|
||||
|
||||
@staticmethod
|
||||
def _get_platform(query, context_vars: dict | None = None) -> str:
|
||||
"""Get platform name from query or context_vars"""
|
||||
# Try WorkflowQuery first
|
||||
if WorkflowMonitoringHelper._is_workflow_query(query) and query.launcher_type:
|
||||
if hasattr(query.launcher_type, 'value'):
|
||||
return query.launcher_type.value
|
||||
return str(query.launcher_type)
|
||||
|
||||
# Fallback to context_vars for launcher_type (person/group)
|
||||
if context_vars and context_vars.get('_launcher_type'):
|
||||
return context_vars['_launcher_type']
|
||||
|
||||
return 'workflow'
|
||||
|
||||
@staticmethod
|
||||
def _get_sender_name(query, context_vars: dict | None = None) -> str | None:
|
||||
"""Get sender name from query or context_vars"""
|
||||
# Try WorkflowQuery first
|
||||
if WorkflowMonitoringHelper._is_workflow_query(query):
|
||||
if query.sender_name:
|
||||
return query.sender_name
|
||||
if query.message_event and hasattr(query.message_event, 'sender'):
|
||||
sender = query.message_event.sender
|
||||
if hasattr(sender, 'nickname'):
|
||||
return sender.nickname
|
||||
if hasattr(sender, 'member_name'):
|
||||
return sender.member_name
|
||||
|
||||
# Fallback to context_vars
|
||||
if context_vars:
|
||||
return context_vars.get('_sender_name')
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def record_trigger_log(
|
||||
ap: app.Application,
|
||||
query,
|
||||
workflow_id: str,
|
||||
workflow_name: str,
|
||||
bot_name: str = 'Workflow',
|
||||
context_vars: dict | None = None,
|
||||
) -> str:
|
||||
"""Record trigger node log (stores original user message content directly)
|
||||
|
||||
Aligned with pipeline monitoring: record_query_start
|
||||
"""
|
||||
try:
|
||||
session_id = WorkflowMonitoringHelper._get_session_id(query, context_vars)
|
||||
platform = WorkflowMonitoringHelper._get_platform(query, context_vars)
|
||||
sender_name = WorkflowMonitoringHelper._get_sender_name(query, context_vars)
|
||||
|
||||
# Get message content - store original content directly
|
||||
message_content = ''
|
||||
if isinstance(query, str):
|
||||
message_content = query
|
||||
elif not isinstance(query, str) and query.message_context:
|
||||
message_content = query.message_context.message_content
|
||||
elif not isinstance(query, str) and query.message_chain and hasattr(query.message_chain, 'model_dump'):
|
||||
message_content = json.dumps(query.message_chain.model_dump(), ensure_ascii=False)
|
||||
elif not isinstance(query, str) and query.user_message:
|
||||
message_content = str(query.user_message)
|
||||
|
||||
# Get bot_id and user_id
|
||||
bot_id = ''
|
||||
user_id = None
|
||||
if not isinstance(query, str):
|
||||
bot_id = query.bot_uuid or ''
|
||||
user_id = query.sender_id
|
||||
elif context_vars:
|
||||
bot_id = context_vars.get('_bot_id', '') or ''
|
||||
user_id = context_vars.get('_user_id')
|
||||
|
||||
message_id = await ap.monitoring_service.record_message(
|
||||
bot_id=bot_id,
|
||||
bot_name=bot_name,
|
||||
pipeline_id=workflow_id,
|
||||
pipeline_name=workflow_name or 'Workflow',
|
||||
message_content=message_content,
|
||||
session_id=session_id,
|
||||
status='success',
|
||||
level='info',
|
||||
platform=platform,
|
||||
user_id=user_id,
|
||||
user_name=sender_name,
|
||||
role='user',
|
||||
runner_name='local-workflow',
|
||||
)
|
||||
|
||||
return message_id
|
||||
except Exception as e:
|
||||
ap.logger.error(f'Failed to record trigger log: {e}')
|
||||
return ''
|
||||
|
||||
@staticmethod
|
||||
async def record_llm_call_log(
|
||||
ap: app.Application,
|
||||
query,
|
||||
workflow_id: str,
|
||||
workflow_name: str,
|
||||
node_name: str,
|
||||
model_name: str,
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
duration_ms: int,
|
||||
status: str = 'success',
|
||||
error_message: str | None = None,
|
||||
bot_name: str = 'Workflow',
|
||||
context_vars: dict | None = None,
|
||||
input_message: str | None = None,
|
||||
message_id: str | None = None,
|
||||
):
|
||||
"""Record LLM call log with message_id association
|
||||
|
||||
Aligned with pipeline monitoring: record_llm_call with message_id
|
||||
LLM calls are aggregated under the trigger log via message_id.
|
||||
"""
|
||||
try:
|
||||
session_id = WorkflowMonitoringHelper._get_session_id(query, context_vars)
|
||||
|
||||
# Get bot_id
|
||||
bot_id = ''
|
||||
if not isinstance(query, str):
|
||||
bot_id = query.bot_uuid or ''
|
||||
elif context_vars:
|
||||
bot_id = context_vars.get('_bot_id', '') or ''
|
||||
|
||||
# Record LLM call with message_id for association
|
||||
await ap.monitoring_service.record_llm_call(
|
||||
bot_id=bot_id,
|
||||
bot_name=bot_name,
|
||||
pipeline_id=workflow_id,
|
||||
pipeline_name=workflow_name or 'Workflow',
|
||||
session_id=session_id,
|
||||
model_name=model_name,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
duration=duration_ms,
|
||||
status=status,
|
||||
error_message=error_message,
|
||||
message_id=message_id,
|
||||
)
|
||||
except Exception as e:
|
||||
ap.logger.error(f'Failed to record LLM call log: {e}')
|
||||
|
||||
@staticmethod
|
||||
async def record_llm_response_log(
|
||||
ap: app.Application,
|
||||
query,
|
||||
workflow_id: str,
|
||||
workflow_name: str,
|
||||
node_name: str,
|
||||
response_content: str,
|
||||
bot_name: str = 'Workflow',
|
||||
context_vars: dict | None = None,
|
||||
):
|
||||
"""Record LLM response log (stores response content directly)
|
||||
|
||||
Aligned with pipeline monitoring: record_query_response
|
||||
"""
|
||||
try:
|
||||
session_id = WorkflowMonitoringHelper._get_session_id(query, context_vars)
|
||||
platform = WorkflowMonitoringHelper._get_platform(query, context_vars)
|
||||
sender_name = WorkflowMonitoringHelper._get_sender_name(query, context_vars)
|
||||
|
||||
# Get bot_id and user_id
|
||||
bot_id = ''
|
||||
user_id = None
|
||||
if not isinstance(query, str):
|
||||
bot_id = query.bot_uuid or ''
|
||||
user_id = query.sender_id
|
||||
elif context_vars:
|
||||
bot_id = context_vars.get('_bot_id', '') or ''
|
||||
user_id = context_vars.get('_user_id')
|
||||
|
||||
# Store response content directly, no prefix
|
||||
await ap.monitoring_service.record_message(
|
||||
bot_id=bot_id,
|
||||
bot_name=bot_name,
|
||||
pipeline_id=workflow_id,
|
||||
pipeline_name=workflow_name or 'Workflow',
|
||||
message_content=response_content[:2000], # Limit length
|
||||
session_id=session_id,
|
||||
status='success',
|
||||
level='info',
|
||||
platform=platform,
|
||||
user_id=user_id,
|
||||
user_name=sender_name,
|
||||
role='assistant',
|
||||
runner_name='local-workflow',
|
||||
)
|
||||
except Exception as e:
|
||||
ap.logger.error(f'Failed to record LLM response log: {e}')
|
||||
|
||||
@staticmethod
|
||||
async def record_reply_log(
|
||||
ap: app.Application,
|
||||
query,
|
||||
workflow_id: str,
|
||||
workflow_name: str,
|
||||
node_name: str,
|
||||
reply_content: str,
|
||||
bot_name: str = 'Workflow',
|
||||
context_vars: dict | None = None,
|
||||
):
|
||||
"""Record reply message log (stores reply content directly)
|
||||
|
||||
Aligned with pipeline monitoring: record_query_response
|
||||
"""
|
||||
try:
|
||||
session_id = WorkflowMonitoringHelper._get_session_id(query, context_vars)
|
||||
platform = WorkflowMonitoringHelper._get_platform(query, context_vars)
|
||||
sender_name = WorkflowMonitoringHelper._get_sender_name(query, context_vars)
|
||||
|
||||
# Get bot_id and user_id
|
||||
bot_id = ''
|
||||
user_id = None
|
||||
if not isinstance(query, str):
|
||||
bot_id = query.bot_uuid or ''
|
||||
user_id = query.sender_id
|
||||
elif context_vars:
|
||||
bot_id = context_vars.get('_bot_id', '') or ''
|
||||
user_id = context_vars.get('_user_id')
|
||||
|
||||
# Store reply content directly, no prefix
|
||||
await ap.monitoring_service.record_message(
|
||||
bot_id=bot_id,
|
||||
bot_name=bot_name,
|
||||
pipeline_id=workflow_id,
|
||||
pipeline_name=workflow_name or 'Workflow',
|
||||
message_content=reply_content[:2000], # Limit length
|
||||
session_id=session_id,
|
||||
status='success',
|
||||
level='info',
|
||||
platform=platform,
|
||||
user_id=user_id,
|
||||
user_name=sender_name,
|
||||
role='assistant',
|
||||
runner_name='local-workflow',
|
||||
)
|
||||
except Exception as e:
|
||||
ap.logger.error(f'Failed to record reply log: {e}')
|
||||
|
||||
|
||||
class LLMCallMonitor:
|
||||
"""Context manager for monitoring LLM calls in workflow"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ap: app.Application,
|
||||
query,
|
||||
bot_id: str,
|
||||
bot_name: str,
|
||||
workflow_id: str,
|
||||
workflow_name: str,
|
||||
node_name: str,
|
||||
model_name: str,
|
||||
context_vars: dict | None = None,
|
||||
):
|
||||
self.ap = ap
|
||||
self.query = query
|
||||
self.bot_id = bot_id
|
||||
self.bot_name = bot_name
|
||||
self.workflow_id = workflow_id
|
||||
self.workflow_name = workflow_name
|
||||
self.node_name = node_name
|
||||
self.model_name = model_name
|
||||
self.context_vars = context_vars
|
||||
self.start_time = None
|
||||
self.input_tokens = 0
|
||||
self.output_tokens = 0
|
||||
|
||||
async def __aenter__(self):
|
||||
self.start_time = time.time()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
duration_ms = int((time.time() - self.start_time) * 1000) if self.start_time else 0
|
||||
|
||||
if exc_type is not None:
|
||||
await WorkflowMonitoringHelper.record_llm_call_log(
|
||||
ap=self.ap,
|
||||
query=self.query,
|
||||
workflow_id=self.workflow_id,
|
||||
workflow_name=self.workflow_name,
|
||||
node_name=self.node_name,
|
||||
model_name=self.model_name,
|
||||
input_tokens=self.input_tokens,
|
||||
output_tokens=self.output_tokens,
|
||||
duration_ms=duration_ms,
|
||||
status='error',
|
||||
error_message=str(exc_val) if exc_val else None,
|
||||
bot_name=self.bot_name,
|
||||
context_vars=self.context_vars,
|
||||
)
|
||||
else:
|
||||
await WorkflowMonitoringHelper.record_llm_call_log(
|
||||
ap=self.ap,
|
||||
query=self.query,
|
||||
workflow_id=self.workflow_id,
|
||||
workflow_name=self.workflow_name,
|
||||
node_name=self.node_name,
|
||||
model_name=self.model_name,
|
||||
input_tokens=self.input_tokens,
|
||||
output_tokens=self.output_tokens,
|
||||
duration_ms=duration_ms,
|
||||
status='success',
|
||||
bot_name=self.bot_name,
|
||||
context_vars=self.context_vars,
|
||||
)
|
||||
|
||||
return False
|
||||
@@ -0,0 +1,164 @@
|
||||
"""Workflow node base class and decorators"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import abc
|
||||
from typing import Any, Callable, Optional, TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .entities import ExecutionContext
|
||||
from ..core import app
|
||||
|
||||
|
||||
class WorkflowNode(abc.ABC):
|
||||
"""Base class for all workflow nodes.
|
||||
|
||||
Node metadata (inputs, outputs, config schema, label, icon, etc.) is
|
||||
defined exclusively in YAML files under templates/metadata/nodes/.
|
||||
Python subclasses only provide execution logic and runtime behaviour.
|
||||
"""
|
||||
|
||||
# Set by @workflow_node decorator
|
||||
type_name: str = ''
|
||||
|
||||
# Category is kept as a fallback for registry when YAML is missing
|
||||
category: str = 'misc'
|
||||
|
||||
# Pipeline config reuse (referenced by registry merge logic)
|
||||
config_schema_source: Optional[str] = None
|
||||
config_stages: list[str] = []
|
||||
|
||||
def __init__(self, node_id: str, config: dict[str, Any], ap: Optional['app.Application'] = None):
|
||||
"""Initialize node with ID and configuration"""
|
||||
self.node_id = node_id
|
||||
self.config = config
|
||||
self.ap = ap
|
||||
|
||||
@abc.abstractmethod
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
"""Execute the node logic.
|
||||
|
||||
Args:
|
||||
inputs: Input data from connected nodes
|
||||
context: Execution context with workflow state
|
||||
|
||||
Returns:
|
||||
Dictionary of output values
|
||||
"""
|
||||
pass
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Validation helpers — metadata is resolved from the registry at
|
||||
# runtime so that YAML remains the single source of truth.
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def validate_inputs(self, inputs: dict[str, Any]) -> list[str]:
|
||||
"""Validate input data against YAML port definitions.
|
||||
|
||||
Returns:
|
||||
List of validation error messages (empty if valid)
|
||||
"""
|
||||
metadata = self._get_metadata()
|
||||
if metadata is None:
|
||||
return []
|
||||
|
||||
errors: list[str] = []
|
||||
for port in metadata.get('inputs', []):
|
||||
if port.get('required', True) and port.get('name') and port['name'] not in inputs:
|
||||
errors.append(f"Missing required input: {port['name']}")
|
||||
return errors
|
||||
|
||||
async def validate_config(self) -> list[str]:
|
||||
"""Validate node configuration against YAML config schema.
|
||||
|
||||
Returns:
|
||||
List of validation error messages (empty if valid)
|
||||
"""
|
||||
metadata = self._get_metadata()
|
||||
if metadata is None:
|
||||
return []
|
||||
|
||||
errors: list[str] = []
|
||||
for cfg in metadata.get('config', []):
|
||||
name = cfg.get('name', '')
|
||||
if not name:
|
||||
continue
|
||||
required = cfg.get('required', False)
|
||||
cfg_type = cfg.get('type', 'string')
|
||||
|
||||
if required and name not in self.config:
|
||||
errors.append(f'Missing required config: {name}')
|
||||
elif name in self.config:
|
||||
value = self.config[name]
|
||||
# Type validation
|
||||
if cfg_type == 'integer' and not isinstance(value, int):
|
||||
errors.append(f'Config {name} must be an integer')
|
||||
elif cfg_type == 'number' and not isinstance(value, (int, float)):
|
||||
errors.append(f'Config {name} must be a number')
|
||||
elif cfg_type == 'boolean' and not isinstance(value, bool):
|
||||
errors.append(f'Config {name} must be a boolean')
|
||||
# Range validation
|
||||
min_val = cfg.get('min_value')
|
||||
max_val = cfg.get('max_value')
|
||||
if min_val is not None and isinstance(value, (int, float)):
|
||||
if value < min_val:
|
||||
errors.append(f'Config {name} must be >= {min_val}')
|
||||
if max_val is not None and isinstance(value, (int, float)):
|
||||
if value > max_val:
|
||||
errors.append(f'Config {name} must be <= {max_val}')
|
||||
return errors
|
||||
|
||||
def get_config(self, key: str, default: Any = None) -> Any:
|
||||
"""Get configuration value with default"""
|
||||
return self.config.get(key, default)
|
||||
|
||||
def _get_metadata(self) -> Optional[dict[str, Any]]:
|
||||
"""Retrieve YAML metadata for this node from the registry."""
|
||||
from .registry import NodeTypeRegistry
|
||||
registry = NodeTypeRegistry.instance()
|
||||
return registry.get_metadata(self.type_name)
|
||||
|
||||
@classmethod
|
||||
def to_schema(cls) -> dict[str, Any]:
|
||||
"""Return a schema dict for this node type.
|
||||
|
||||
This is used by tests and tooling to inspect node capabilities.
|
||||
"""
|
||||
from .registry import NodeTypeRegistry
|
||||
registry = NodeTypeRegistry.instance()
|
||||
metadata = registry.get_metadata(cls.type_name)
|
||||
if metadata:
|
||||
return registry._metadata_to_schema(metadata)
|
||||
# Fallback: build a minimal schema from class attributes
|
||||
return {
|
||||
'type': f'{cls.category}.{cls.type_name}' if cls.type_name else cls.type_name,
|
||||
'category': cls.category,
|
||||
'label': getattr(cls, 'name', cls.type_name),
|
||||
'description': getattr(cls, 'description', ''),
|
||||
'inputs': [],
|
||||
'outputs': [],
|
||||
'config_schema': [],
|
||||
}
|
||||
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Decorator for setting type_name attribute
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
|
||||
def workflow_node(type_name: str) -> Callable[[type[WorkflowNode]], type[WorkflowNode]]:
|
||||
"""Decorator to set the type_name attribute on a workflow node class.
|
||||
|
||||
Usage:
|
||||
@workflow_node('llm_call')
|
||||
class LLMCallNode(WorkflowNode):
|
||||
...
|
||||
|
||||
The actual registration is now handled by the discovery engine.
|
||||
"""
|
||||
|
||||
def decorator(cls: type[WorkflowNode]) -> type[WorkflowNode]:
|
||||
cls.type_name = type_name
|
||||
return cls
|
||||
|
||||
return decorator
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,263 @@
|
||||
"""Call Pipeline Node - invoke an existing pipeline
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/call_pipeline.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Optional
|
||||
|
||||
import pydantic
|
||||
|
||||
import langbot_plugin.api.definition.abstract.platform.adapter as abstract_platform_adapter
|
||||
import langbot_plugin.api.definition.abstract.platform.event_logger as abstract_event_logger
|
||||
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
|
||||
import langbot_plugin.api.entities.builtin.platform.entities as platform_entities
|
||||
import langbot_plugin.api.entities.builtin.platform.events as platform_events
|
||||
import langbot_plugin.api.entities.builtin.platform.message as platform_message
|
||||
import langbot_plugin.api.entities.builtin.provider.session as provider_session
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
|
||||
class _NoOpEventLogger(abstract_event_logger.AbstractEventLogger):
|
||||
"""No-op event logger for workflow pipeline adapter."""
|
||||
|
||||
async def info(
|
||||
self,
|
||||
text: str,
|
||||
images: Optional[list[platform_message.Image]] = None,
|
||||
message_session_id: Optional[str] = None,
|
||||
no_throw: bool = True,
|
||||
):
|
||||
pass
|
||||
|
||||
async def debug(
|
||||
self,
|
||||
text: str,
|
||||
images: Optional[list[platform_message.Image]] = None,
|
||||
message_session_id: Optional[str] = None,
|
||||
no_throw: bool = True,
|
||||
):
|
||||
pass
|
||||
|
||||
async def warning(
|
||||
self,
|
||||
text: str,
|
||||
images: Optional[list[platform_message.Image]] = None,
|
||||
message_session_id: Optional[str] = None,
|
||||
no_throw: bool = True,
|
||||
):
|
||||
pass
|
||||
|
||||
async def error(
|
||||
self,
|
||||
text: str,
|
||||
images: Optional[list[platform_message.Image]] = None,
|
||||
message_session_id: Optional[str] = None,
|
||||
no_throw: bool = True,
|
||||
):
|
||||
pass
|
||||
|
||||
@workflow_node('call_pipeline')
|
||||
class CallPipelineNode(WorkflowNode):
|
||||
"""Call pipeline node - invoke an existing pipeline"""
|
||||
|
||||
category = 'action'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
if not self.ap:
|
||||
raise RuntimeError('Application instance not available — cannot call pipeline')
|
||||
|
||||
raw_query = inputs.get('query', '')
|
||||
query_text = str(raw_query or inputs.get('input') or '')
|
||||
pipeline_ref = str(self.get_config('pipeline_uuid', '') or '').strip()
|
||||
|
||||
if not pipeline_ref:
|
||||
raise ValueError('No pipeline configured for call pipeline node')
|
||||
|
||||
pipeline_data = await self.ap.pipeline_service.get_pipeline(pipeline_ref)
|
||||
if pipeline_data is None:
|
||||
pipeline_data = await self.ap.pipeline_service.get_pipeline_by_name(pipeline_ref)
|
||||
if pipeline_data is None:
|
||||
raise ValueError(f'Pipeline not found: {pipeline_ref}')
|
||||
|
||||
pipeline_uuid = str(pipeline_data.get('uuid', '') or '')
|
||||
if not pipeline_uuid:
|
||||
raise ValueError(f'Pipeline UUID missing for: {pipeline_ref}')
|
||||
|
||||
runtime_pipeline = await self.ap.pipeline_mgr.get_pipeline_by_uuid(pipeline_uuid)
|
||||
if runtime_pipeline is None:
|
||||
raise ValueError(f'Runtime pipeline not loaded: {pipeline_uuid}')
|
||||
|
||||
adapter = _WorkflowPipelineCaptureAdapter(context=context)
|
||||
adapter.bot_account_id = 'workflow-call-pipeline'
|
||||
|
||||
message_event = self._build_message_event(query_text, context)
|
||||
message_chain = message_event.message_chain
|
||||
launcher_type = (
|
||||
provider_session.LauncherTypes.GROUP
|
||||
if context.message_context and context.message_context.is_group
|
||||
else provider_session.LauncherTypes.PERSON
|
||||
)
|
||||
launcher_id = context.session_id or context.execution_id
|
||||
sender_id = (
|
||||
context.message_context.sender_id
|
||||
if context.message_context and context.message_context.sender_id
|
||||
else context.user_id or f'workflow_{context.execution_id}'
|
||||
)
|
||||
|
||||
query = pipeline_query.Query(
|
||||
bot_uuid=context.bot_id,
|
||||
query_id=-1,
|
||||
launcher_type=launcher_type,
|
||||
launcher_id=launcher_id,
|
||||
sender_id=sender_id,
|
||||
message_event=message_event,
|
||||
message_chain=message_chain,
|
||||
variables={
|
||||
'_called_from_workflow': True,
|
||||
'_workflow_execution_id': context.execution_id,
|
||||
'_workflow_id': context.workflow_id,
|
||||
**dict(context.variables or {}),
|
||||
},
|
||||
resp_messages=[],
|
||||
resp_message_chain=[],
|
||||
adapter=adapter,
|
||||
pipeline_uuid=pipeline_uuid,
|
||||
)
|
||||
|
||||
await runtime_pipeline.run(query)
|
||||
|
||||
response_text = adapter.get_last_text_response()
|
||||
result = {
|
||||
'pipeline_uuid': pipeline_uuid,
|
||||
'pipeline_name': pipeline_data.get('name', ''),
|
||||
'responses': adapter.responses,
|
||||
'query_text': query_text,
|
||||
}
|
||||
|
||||
return {'response': response_text, 'result': result}
|
||||
|
||||
def _build_message_event(
|
||||
self,
|
||||
query_text: str,
|
||||
context: ExecutionContext,
|
||||
) -> platform_events.MessageEvent:
|
||||
message_chain_data = context.trigger_data.get('message_chain') or context.trigger_data.get('message', [])
|
||||
if isinstance(message_chain_data, list) and message_chain_data:
|
||||
message_chain = platform_message.MessageChain.model_validate(message_chain_data)
|
||||
else:
|
||||
message_chain = platform_message.MessageChain([platform_message.Plain(text=query_text)])
|
||||
|
||||
if context.message_context and context.message_context.is_group:
|
||||
group = platform_entities.Group(
|
||||
id=context.message_context.group_id or context.session_id or 'workflow_group',
|
||||
name=context.message_context.raw_message.get('group_name', 'Workflow Group') if context.message_context.raw_message else 'Workflow Group',
|
||||
permission=platform_entities.Permission.Member,
|
||||
)
|
||||
sender = platform_entities.GroupMember(
|
||||
id=context.message_context.sender_id,
|
||||
member_name=context.message_context.sender_name or 'Workflow User',
|
||||
permission=platform_entities.Permission.Member,
|
||||
group=group,
|
||||
)
|
||||
return platform_events.GroupMessage(
|
||||
sender=sender,
|
||||
message_chain=message_chain,
|
||||
time=context.message_context.raw_message.get('time') if context.message_context.raw_message else None,
|
||||
)
|
||||
|
||||
sender = platform_entities.Friend(
|
||||
id=context.message_context.sender_id if context.message_context else context.user_id or 'workflow_user',
|
||||
nickname=context.message_context.sender_name if context.message_context else 'Workflow User',
|
||||
remark=context.message_context.sender_name if context.message_context else 'Workflow User',
|
||||
)
|
||||
return platform_events.FriendMessage(
|
||||
sender=sender,
|
||||
message_chain=message_chain,
|
||||
time=context.message_context.raw_message.get('time')
|
||||
if context.message_context and context.message_context.raw_message
|
||||
else None,
|
||||
)
|
||||
|
||||
class _WorkflowPipelineCaptureAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
|
||||
"""Adapter to capture pipeline responses for workflow execution."""
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
responses: list[dict[str, Any]] = []
|
||||
context: Optional[ExecutionContext] = pydantic.Field(default=None, exclude=True)
|
||||
|
||||
def __init__(self, context: ExecutionContext):
|
||||
super().__init__(config={}, logger=_NoOpEventLogger(), context=context)
|
||||
self.responses = []
|
||||
|
||||
async def send_message(self, target_type: str, target_id: str, message: platform_message.MessageChain):
|
||||
payload = {
|
||||
'type': 'send',
|
||||
'target_type': target_type,
|
||||
'target_id': target_id,
|
||||
'content': str(message),
|
||||
'message_chain': message.model_dump(),
|
||||
}
|
||||
self.responses.append(payload)
|
||||
return payload
|
||||
|
||||
async def reply_message(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
message: platform_message.MessageChain,
|
||||
quote_origin: bool = False,
|
||||
):
|
||||
payload = {
|
||||
'type': 'reply',
|
||||
'content': str(message),
|
||||
'message_chain': message.model_dump(),
|
||||
'quote_origin': quote_origin,
|
||||
}
|
||||
self.responses.append(payload)
|
||||
return payload
|
||||
|
||||
async def reply_message_chunk(
|
||||
self,
|
||||
message_source: platform_events.MessageEvent,
|
||||
bot_message: dict,
|
||||
message: platform_message.MessageChain,
|
||||
quote_origin: bool = False,
|
||||
is_final: bool = False,
|
||||
):
|
||||
payload = {
|
||||
'type': 'reply_chunk',
|
||||
'content': str(message),
|
||||
'message_chain': message.model_dump(),
|
||||
'quote_origin': quote_origin,
|
||||
'is_final': is_final,
|
||||
}
|
||||
self.responses.append(payload)
|
||||
return payload
|
||||
|
||||
async def create_message_card(self, message_id, event: platform_events.MessageEvent) -> bool:
|
||||
return False
|
||||
|
||||
def register_listener(self, event_type, callback):
|
||||
return None
|
||||
|
||||
def unregister_listener(self, event_type, callback):
|
||||
return None
|
||||
|
||||
async def run_async(self):
|
||||
return None
|
||||
|
||||
async def is_stream_output_supported(self) -> bool:
|
||||
return False
|
||||
|
||||
async def kill(self) -> bool:
|
||||
return True
|
||||
|
||||
def get_last_text_response(self) -> str:
|
||||
if not self.responses:
|
||||
return ''
|
||||
return str(self.responses[-1].get('content', '') or '')
|
||||
@@ -0,0 +1,85 @@
|
||||
"""Call Workflow Node - invoke an existing workflow
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/call_workflow.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
|
||||
@workflow_node('call_workflow')
|
||||
class CallWorkflowNode(WorkflowNode):
|
||||
"""Call workflow node - invoke an existing workflow"""
|
||||
|
||||
category = 'action'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
if not self.ap:
|
||||
raise RuntimeError('Application instance not available — cannot call workflow')
|
||||
|
||||
# Get workflow reference from config
|
||||
workflow_ref = str(self.get_config('workflow_uuid', '') or '').strip()
|
||||
if not workflow_ref:
|
||||
raise ValueError('No workflow configured for call workflow node')
|
||||
|
||||
# Get workflow definition from service
|
||||
workflow_data = await self.ap.workflow_service.get_workflow(workflow_ref)
|
||||
if workflow_data is None:
|
||||
raise ValueError(f'Workflow not found: {workflow_ref}')
|
||||
|
||||
workflow_uuid = str(workflow_data.get('uuid', '') or '')
|
||||
if not workflow_uuid:
|
||||
raise ValueError(f'Workflow UUID missing for: {workflow_ref}')
|
||||
|
||||
# Build variables to pass to the called workflow
|
||||
variables = dict(inputs.get('variables', {}) or {})
|
||||
|
||||
# Inherit current workflow variables if configured
|
||||
if self.get_config('inherit_variables', True):
|
||||
for key, value in (context.variables or {}).items():
|
||||
if key not in variables:
|
||||
variables[key] = value
|
||||
|
||||
# Add context markers for debugging
|
||||
variables['_called_from_workflow'] = True
|
||||
variables['_parent_workflow_id'] = context.workflow_id
|
||||
variables['_parent_execution_id'] = context.execution_id
|
||||
|
||||
# Execute the workflow
|
||||
execution_id = await self.ap.workflow_service.execute_workflow(
|
||||
workflow_uuid=workflow_uuid,
|
||||
trigger_type='workflow_call',
|
||||
trigger_data={
|
||||
'variables': variables,
|
||||
'parent_execution_id': context.execution_id,
|
||||
},
|
||||
session_id=context.session_id,
|
||||
user_id=context.user_id,
|
||||
bot_id=context.bot_id,
|
||||
)
|
||||
|
||||
# Get execution result
|
||||
execution = await self.ap.workflow_service.get_execution(execution_id)
|
||||
if execution is None:
|
||||
raise ValueError(f'Execution result not found: {execution_id}')
|
||||
|
||||
# Build result
|
||||
result = {
|
||||
'workflow_uuid': workflow_uuid,
|
||||
'workflow_name': workflow_data.get('name', ''),
|
||||
'execution_id': execution_id,
|
||||
'status': execution.get('status', 'unknown'),
|
||||
'variables': execution.get('variables', {}),
|
||||
'error': execution.get('error'),
|
||||
}
|
||||
|
||||
return {
|
||||
'result': result,
|
||||
'status': execution.get('status', 'unknown'),
|
||||
'error': execution.get('error'),
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
"""Code Executor Node - run Python or JavaScript code
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/code_executor.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import io
|
||||
import logging
|
||||
import sys
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 危险的内置函数和模块黑名单
|
||||
_DANGEROUS_BUILTINS = {
|
||||
'__import__', 'eval', 'exec', 'compile', 'open', 'file',
|
||||
'input', 'exit', 'quit', 'globals', 'locals', 'vars',
|
||||
'dir', 'help', 'breakpoint',
|
||||
}
|
||||
|
||||
# 允许的安全内置函数
|
||||
_SAFE_BUILTINS = {
|
||||
'abs': abs, 'all': all, 'any': any, 'bin': bin, 'bool': bool,
|
||||
'bytearray': bytearray, 'bytes': bytes, 'callable': callable,
|
||||
'chr': chr, 'complex': complex, 'dict': dict, 'divmod': divmod,
|
||||
'enumerate': enumerate, 'filter': filter, 'float': float,
|
||||
'format': format, 'frozenset': frozenset, 'hash': hash,
|
||||
'hex': hex, 'int': int, 'isinstance': isinstance, 'issubclass': issubclass,
|
||||
'iter': iter, 'len': len, 'list': list, 'map': map, 'max': max,
|
||||
'min': min, 'next': next, 'object': object, 'oct': oct, 'ord': ord,
|
||||
'pow': pow, 'print': print, 'range': range, 'repr': repr,
|
||||
'reversed': reversed, 'round': round, 'set': set, 'slice': slice,
|
||||
'sorted': sorted, 'str': str, 'sum': sum, 'tuple': tuple,
|
||||
'type': type, 'zip': zip,
|
||||
}
|
||||
|
||||
|
||||
def _check_code_safety(code: str) -> list[str]:
|
||||
"""检查代码中是否包含危险操作"""
|
||||
warnings = []
|
||||
try:
|
||||
tree = ast.parse(code)
|
||||
for node in ast.walk(tree):
|
||||
# 检查 import 语句
|
||||
if isinstance(node, (ast.Import, ast.ImportFrom)):
|
||||
warnings.append('Import statements are not allowed')
|
||||
# 检查危险函数调用
|
||||
if isinstance(node, ast.Call):
|
||||
if isinstance(node.func, ast.Name) and node.func.id in _DANGEROUS_BUILTINS:
|
||||
warnings.append(f'Dangerous function call: {node.func.id}')
|
||||
# 检查 __import__ 通过 getattr 调用
|
||||
if isinstance(node.func, ast.Attribute):
|
||||
if node.func.attr in ('__import__', 'eval', 'exec', 'open', 'file'):
|
||||
warnings.append(f'Dangerous attribute access: {node.func.attr}')
|
||||
except SyntaxError as e:
|
||||
warnings.append(f'Syntax error in code: {e}')
|
||||
return warnings
|
||||
|
||||
|
||||
class _ExecutionTimeoutError(Exception):
|
||||
"""执行超时错误"""
|
||||
pass
|
||||
|
||||
|
||||
def _run_with_timeout(func, timeout: float = 10.0):
|
||||
"""带超时限制的函数执行"""
|
||||
result = [None]
|
||||
error = [None]
|
||||
|
||||
def _target():
|
||||
try:
|
||||
result[0] = func()
|
||||
except Exception as e:
|
||||
error[0] = e
|
||||
|
||||
thread = threading.Thread(target=_target)
|
||||
thread.daemon = True
|
||||
thread.start()
|
||||
thread.join(timeout)
|
||||
|
||||
if thread.is_alive():
|
||||
raise _ExecutionTimeoutError(f'Code execution timed out after {timeout} seconds')
|
||||
|
||||
if error[0]:
|
||||
raise error[0]
|
||||
|
||||
return result[0]
|
||||
|
||||
|
||||
@workflow_node('code_executor')
|
||||
class CodeExecutorNode(WorkflowNode):
|
||||
"""Code executor node - run Python or JavaScript code"""
|
||||
|
||||
category = 'process'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
code = self.get_config('code', '')
|
||||
language = self.get_config('language', 'python')
|
||||
timeout = self.get_config('timeout', 10)
|
||||
|
||||
# 限制最大超时时间
|
||||
timeout = min(max(timeout, 1), 30)
|
||||
|
||||
if not code:
|
||||
return {'output': None, 'console': '', 'error': 'No code provided'}
|
||||
|
||||
if language == 'python':
|
||||
return await self._execute_python(code, inputs, context, timeout)
|
||||
else:
|
||||
return await self._execute_javascript(code, inputs, context)
|
||||
|
||||
async def _execute_python(self, code: str, inputs: dict[str, Any], context: ExecutionContext, timeout: float) -> dict[str, Any]:
|
||||
# 安全检查
|
||||
warnings = _check_code_safety(code)
|
||||
if warnings:
|
||||
logger.warning('Code safety warnings: %s', warnings)
|
||||
return {'output': None, 'console': '', 'error': '; '.join(warnings)}
|
||||
|
||||
stdout_capture = io.StringIO()
|
||||
old_stdout = sys.stdout
|
||||
|
||||
def _exec_code():
|
||||
nonlocal stdout_capture
|
||||
sys.stdout = stdout_capture
|
||||
try:
|
||||
# 使用更安全的执行方式
|
||||
compiled = compile(code, '<workflow>', 'exec')
|
||||
safe_globals = {
|
||||
'__builtins__': _SAFE_BUILTINS,
|
||||
'__name__': '__workflow_sandbox__',
|
||||
}
|
||||
local_vars = {'inputs': inputs, 'output': None}
|
||||
exec(compiled, safe_globals, local_vars)
|
||||
return local_vars.get('output')
|
||||
finally:
|
||||
sys.stdout = old_stdout
|
||||
|
||||
try:
|
||||
output = _run_with_timeout(_exec_code, timeout)
|
||||
console_output = stdout_capture.getvalue()
|
||||
return {'output': output, 'console': console_output, 'error': None}
|
||||
except _ExecutionTimeoutError as e:
|
||||
logger.error('Code execution timeout: %s', e)
|
||||
return {'output': None, 'console': stdout_capture.getvalue(), 'error': str(e)}
|
||||
except Exception as e:
|
||||
logger.error('Code execution error: %s', e)
|
||||
return {'output': None, 'console': stdout_capture.getvalue(), 'error': f'{type(e).__name__}: {e}'}
|
||||
|
||||
async def _execute_javascript(self, code: str, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
return {'output': None, 'console': '', 'error': 'JavaScript execution is not implemented'}
|
||||
@@ -0,0 +1,125 @@
|
||||
"""Condition Node - branch based on condition
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/condition.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import signal
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
from ..safe_eval import safe_eval_with_vars
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 正则表达式超时限制(秒)
|
||||
_REGEX_TIMEOUT = 2
|
||||
|
||||
|
||||
class _RegexTimeoutError(Exception):
|
||||
"""正则表达式超时错误"""
|
||||
pass
|
||||
|
||||
|
||||
def _handle_timeout(signum, frame):
|
||||
"""超时信号处理"""
|
||||
raise _RegexTimeoutError('Regex match timed out')
|
||||
|
||||
|
||||
def _safe_regex_match(pattern: str, text: str) -> tuple[bool, str]:
|
||||
"""安全地执行正则表达式匹配,带有超时限制"""
|
||||
# 设置超时信号
|
||||
old_handler = signal.signal(signal.SIGALRM, _handle_timeout)
|
||||
signal.setitimer(signal.ITIMER_REAL, _REGEX_TIMEOUT)
|
||||
|
||||
try:
|
||||
result = bool(re.match(pattern, str(text)))
|
||||
return result, ''
|
||||
except _RegexTimeoutError:
|
||||
logger.warning('Regex match timed out for pattern: %s', pattern[:50])
|
||||
return False, 'Regex match timed out'
|
||||
except re.error as e:
|
||||
logger.warning('Invalid regex pattern: %s', e)
|
||||
return False, f'Invalid regex: {e}'
|
||||
finally:
|
||||
signal.setitimer(signal.ITIMER_REAL, 0)
|
||||
signal.signal(signal.SIGALRM, old_handler)
|
||||
|
||||
|
||||
@workflow_node('condition')
|
||||
class ConditionNode(WorkflowNode):
|
||||
"""Condition node - branch based on condition"""
|
||||
|
||||
category = 'control'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
condition_type = self.get_config('condition_type', 'expression')
|
||||
input_data = inputs.get('input')
|
||||
|
||||
result = False
|
||||
|
||||
if condition_type == 'expression':
|
||||
expression = self.get_config('expression', 'false')
|
||||
result = await self._evaluate_expression(expression, input_data, context)
|
||||
elif condition_type == 'comparison':
|
||||
result = await self._evaluate_comparison(input_data, context)
|
||||
elif condition_type == 'contains':
|
||||
left = self.get_config('left_value', '')
|
||||
right = self.get_config('right_value', '')
|
||||
result = right in left
|
||||
elif condition_type == 'empty':
|
||||
result = not bool(input_data)
|
||||
elif condition_type == 'regex':
|
||||
left = self.get_config('left_value', '')
|
||||
pattern = self.get_config('right_value', '')
|
||||
result, error = _safe_regex_match(pattern, left)
|
||||
if error:
|
||||
return {'true': None, 'false': input_data, 'error': error}
|
||||
|
||||
if result:
|
||||
return {'true': input_data, 'false': None}
|
||||
else:
|
||||
return {'true': None, 'false': input_data}
|
||||
|
||||
async def _evaluate_expression(self, expression: str, data: Any, context: ExecutionContext) -> bool:
|
||||
try:
|
||||
local_vars = {'input': data, 'data': data, 'variables': context.variables}
|
||||
return bool(safe_eval_with_vars(expression, local_vars))
|
||||
except Exception as e:
|
||||
logger.warning('Expression evaluation error: %s', e)
|
||||
return False
|
||||
|
||||
async def _evaluate_comparison(self, data: Any, context: ExecutionContext) -> bool:
|
||||
left = self.get_config('left_value', '')
|
||||
right = self.get_config('right_value', '')
|
||||
operator = self.get_config('operator', '==')
|
||||
|
||||
try:
|
||||
left_num = float(left)
|
||||
right_num = float(right)
|
||||
|
||||
if operator == '==':
|
||||
return left_num == right_num
|
||||
elif operator == '!=':
|
||||
return left_num != right_num
|
||||
elif operator == '>':
|
||||
return left_num > right_num
|
||||
elif operator == '<':
|
||||
return left_num < right_num
|
||||
elif operator == '>=':
|
||||
return left_num >= right_num
|
||||
elif operator == '<=':
|
||||
return left_num <= right_num
|
||||
except ValueError:
|
||||
if operator == '==':
|
||||
return left == right
|
||||
elif operator == '!=':
|
||||
return left != right
|
||||
elif operator in ('>', '<', '>=', '<='):
|
||||
return False
|
||||
|
||||
return False
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Coze Bot Node - call Coze API bot
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/coze_bot.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('coze_bot')
|
||||
class CozeBotNode(WorkflowNode):
|
||||
"""Coze bot node - call Coze API bot"""
|
||||
|
||||
category = 'integration'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
api_key = self.get_config('api_key', '')
|
||||
bot_id = self.get_config('bot_id', '')
|
||||
api_base = self.get_config('api_base', 'https://api.coze.cn')
|
||||
query = inputs.get('query', '')
|
||||
conversation_id = inputs.get('conversation_id')
|
||||
|
||||
# Safe API key truncation
|
||||
masked_key = f'{api_key[:4]}...{api_key[-4:]}' if len(api_key) > 8 else '***' if api_key else ''
|
||||
|
||||
return {
|
||||
'answer': '',
|
||||
'conversation_id': conversation_id,
|
||||
'success': False,
|
||||
'_debug': {
|
||||
'api_key': masked_key,
|
||||
'bot_id': bot_id,
|
||||
'api_base': api_base,
|
||||
'query': query,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
"""Cron Trigger Node - triggers workflow on schedule
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/cron_trigger.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('cron_trigger')
|
||||
class CronTriggerNode(WorkflowNode):
|
||||
"""Cron trigger node - triggers workflow on schedule"""
|
||||
|
||||
category = 'trigger'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
from datetime import datetime
|
||||
|
||||
return {
|
||||
'timestamp': datetime.now().isoformat(),
|
||||
'schedule': self.get_config('cron', ''),
|
||||
'context': context.trigger_data,
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Data Transform Node - transform data using templates or JSONPath
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/data_transform.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
from ..safe_eval import safe_eval_with_vars
|
||||
|
||||
@workflow_node('data_transform')
|
||||
class DataTransformNode(WorkflowNode):
|
||||
"""Data transform node - transform data using templates or JSONPath"""
|
||||
|
||||
category = 'process'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
data = inputs.get('data')
|
||||
transform_type = self.get_config('transform_type', 'template')
|
||||
|
||||
if transform_type == 'template':
|
||||
template = self.get_config('template', '')
|
||||
result = self._apply_template(template, data, context)
|
||||
elif transform_type == 'jsonpath':
|
||||
expression = self.get_config('expression', '$')
|
||||
result = self._apply_jsonpath(expression, data)
|
||||
elif transform_type == 'expression':
|
||||
expression = self.get_config('expression', '')
|
||||
result = self._evaluate_expression(expression, data, context)
|
||||
else:
|
||||
result = data
|
||||
|
||||
return {'result': result}
|
||||
|
||||
def _apply_template(self, template: str, data: Any, context: ExecutionContext) -> str:
|
||||
result = template
|
||||
if isinstance(data, dict):
|
||||
for key, value in data.items():
|
||||
result = result.replace(f'{{{{data.{key}}}}}', str(value))
|
||||
for key, value in context.variables.items():
|
||||
result = result.replace(f'{{{{variables.{key}}}}}', str(value))
|
||||
return result
|
||||
|
||||
def _apply_jsonpath(self, expression: str, data: Any) -> Any:
|
||||
if expression == '$':
|
||||
return data
|
||||
if expression.startswith('$.'):
|
||||
parts = expression[2:].split('.')
|
||||
result = data
|
||||
for part in parts:
|
||||
if isinstance(result, dict):
|
||||
result = result.get(part)
|
||||
elif isinstance(result, list) and part.isdigit():
|
||||
result = result[int(part)]
|
||||
else:
|
||||
return None
|
||||
return result
|
||||
return data
|
||||
|
||||
def _evaluate_expression(self, expression: str, data: Any, context: ExecutionContext) -> Any:
|
||||
local_vars = {'data': data, 'variables': context.variables}
|
||||
try:
|
||||
return safe_eval_with_vars(expression, local_vars)
|
||||
except Exception:
|
||||
return None
|
||||
@@ -0,0 +1,38 @@
|
||||
"""Database Query Node - execute database queries
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/database_query.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('database_query')
|
||||
class DatabaseQueryNode(WorkflowNode):
|
||||
"""Database query node - execute database queries"""
|
||||
|
||||
category = 'integration'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
connection_type = self.get_config('connection_type', 'postgresql')
|
||||
query = self.get_config('query', '')
|
||||
query_type = self.get_config('query_type', 'select')
|
||||
timeout = self.get_config('timeout', 30)
|
||||
|
||||
parameters = inputs.get('parameters', {})
|
||||
|
||||
return {
|
||||
'results': [],
|
||||
'row_count': 0,
|
||||
'success': False,
|
||||
'_debug': {
|
||||
'connection_type': connection_type,
|
||||
'query': query,
|
||||
'query_type': query_type,
|
||||
'timeout': timeout,
|
||||
'parameters': parameters,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
"""Dify Knowledge Query Node - query Dify knowledge base
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/dify_knowledge_query.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('dify_knowledge_query')
|
||||
class DifyKnowledgeQueryNode(WorkflowNode):
|
||||
"""Dify knowledge base query node - query Dify knowledge base"""
|
||||
|
||||
category = 'integration'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
base_url = self.get_config('base_url', 'https://api.dify.ai/v1')
|
||||
api_key = self.get_config('api_key', '')
|
||||
dataset_id = self.get_config('dataset_id', '')
|
||||
query = inputs.get('query', '')
|
||||
|
||||
# Safe API key truncation
|
||||
masked_key = f'{api_key[:4]}...{api_key[-4:]}' if len(api_key) > 8 else '***' if api_key else ''
|
||||
|
||||
return {
|
||||
'results': [],
|
||||
'success': False,
|
||||
'_debug': {
|
||||
'base_url': base_url,
|
||||
'api_key': masked_key,
|
||||
'dataset_id': dataset_id,
|
||||
'query': query,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Dify Workflow Node - call Dify service API
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/dify_workflow.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('dify_workflow')
|
||||
class DifyWorkflowNode(WorkflowNode):
|
||||
"""Dify workflow node - call Dify service API"""
|
||||
|
||||
category = 'integration'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
base_url = self.get_config('base_url', 'https://api.dify.ai/v1')
|
||||
api_key = self.get_config('api_key', '')
|
||||
app_type = self.get_config('app_type', 'chat')
|
||||
query = inputs.get('query', '')
|
||||
conversation_id = inputs.get('conversation_id')
|
||||
|
||||
# Safe API key truncation
|
||||
masked_key = f'{api_key[:4]}...{api_key[-4:]}' if len(api_key) > 8 else '***' if api_key else ''
|
||||
|
||||
return {
|
||||
'answer': '',
|
||||
'conversation_id': conversation_id,
|
||||
'success': False,
|
||||
'_debug': {
|
||||
'base_url': base_url,
|
||||
'api_key': masked_key,
|
||||
'app_type': app_type,
|
||||
'query': query,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
"""End Node - marks the end of workflow execution
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/end.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('end')
|
||||
class EndNode(WorkflowNode):
|
||||
"""End node - marks the end of workflow execution"""
|
||||
|
||||
category = 'control'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
result = inputs.get('result')
|
||||
output_format = self.get_config('output_format', 'passthrough')
|
||||
|
||||
if output_format == 'text':
|
||||
return {'output': str(result)}
|
||||
elif output_format == 'json':
|
||||
import json
|
||||
|
||||
try:
|
||||
return {'output': json.dumps(result, ensure_ascii=False)}
|
||||
except Exception:
|
||||
return {'output': str(result)}
|
||||
else:
|
||||
return {'output': result}
|
||||
@@ -0,0 +1,28 @@
|
||||
"""Event Trigger Node - triggers workflow on system events
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/event_trigger.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('event_trigger')
|
||||
class EventTriggerNode(WorkflowNode):
|
||||
"""Event trigger node - triggers workflow on system events"""
|
||||
|
||||
category = 'trigger'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
# Safe access to trigger_data which may be None
|
||||
trigger_data = context.trigger_data or {}
|
||||
|
||||
return {
|
||||
'event_type': trigger_data.get('event_type', ''),
|
||||
'event_data': trigger_data.get('event_data', {}),
|
||||
'timestamp': trigger_data.get('timestamp', datetime.now().isoformat()),
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
"""HTTP Request Node - make HTTP API calls
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/http_request.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import logging
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 内网地址黑名单
|
||||
_PRIVATE_NETWORKS = [
|
||||
ipaddress.ip_network('10.0.0.0/8'),
|
||||
ipaddress.ip_network('172.16.0.0/12'),
|
||||
ipaddress.ip_network('192.168.0.0/16'),
|
||||
ipaddress.ip_network('127.0.0.0/8'),
|
||||
ipaddress.ip_network('169.254.0.0/16'),
|
||||
ipaddress.ip_network('0.0.0.0/8'),
|
||||
ipaddress.ip_network('::1/128'),
|
||||
ipaddress.ip_network('fc00::/7'),
|
||||
ipaddress.ip_network('fe80::/10'),
|
||||
]
|
||||
|
||||
# 危险协议
|
||||
_DANGEROUS_SCHEMES = {'file', 'gopher', 'dict', 'ftp', 'telnet'}
|
||||
|
||||
|
||||
def _is_safe_url(url: str) -> tuple[bool, str]:
|
||||
"""检查 URL 是否安全(非内网地址)"""
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except Exception as e:
|
||||
return False, f'Invalid URL: {e}'
|
||||
|
||||
# 检查协议
|
||||
scheme = parsed.scheme.lower()
|
||||
if scheme in _DANGEROUS_SCHEMES:
|
||||
return False, f'Dangerous scheme: {scheme}'
|
||||
|
||||
if scheme not in ('http', 'https'):
|
||||
return False, f'Unsupported scheme: {scheme}'
|
||||
|
||||
# 检查主机名
|
||||
hostname = parsed.hostname
|
||||
if not hostname:
|
||||
return False, 'Missing hostname'
|
||||
|
||||
# 检查是否是危险主机名
|
||||
dangerous_hosts = {'localhost', '0.0.0.0', '127.0.0.1', '::1'}
|
||||
if hostname.lower() in dangerous_hosts:
|
||||
return False, f'Dangerous hostname: {hostname}'
|
||||
|
||||
# 解析 IP 地址并检查是否在私有网络
|
||||
try:
|
||||
ip = ipaddress.ip_address(hostname)
|
||||
for network in _PRIVATE_NETWORKS:
|
||||
if ip in network:
|
||||
return False, f'Private network address: {ip}'
|
||||
except ValueError:
|
||||
# 不是 IP 地址,尝试 DNS 解析检查
|
||||
# 这里可以添加 DNS 解析检查,但为了避免复杂性,暂时跳过
|
||||
pass
|
||||
|
||||
return True, ''
|
||||
|
||||
|
||||
@workflow_node('http_request')
|
||||
class HTTPRequestNode(WorkflowNode):
|
||||
"""HTTP request node - make HTTP API calls"""
|
||||
|
||||
category = 'action'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
import aiohttp
|
||||
|
||||
url = self.get_config('url', '')
|
||||
method = self.get_config('method', 'GET').upper()
|
||||
timeout = self.get_config('timeout', 30)
|
||||
content_type = self.get_config('content_type', 'application/json')
|
||||
allow_redirects = self.get_config('allow_redirects', False) # 默认禁用重定向
|
||||
|
||||
# 限制超时时间
|
||||
timeout = min(max(timeout, 1), 120)
|
||||
|
||||
if not url:
|
||||
return {'response': None, 'status_code': 0, 'headers': {}, 'error': 'No URL provided'}
|
||||
|
||||
# 安全检查 URL
|
||||
is_safe, error_msg = _is_safe_url(url)
|
||||
if not is_safe:
|
||||
logger.warning('Unsafe URL blocked: %s - %s', url, error_msg)
|
||||
return {'response': None, 'status_code': 0, 'headers': {}, 'error': f'Unsafe URL: {error_msg}'}
|
||||
|
||||
# 验证 HTTP 方法
|
||||
allowed_methods = {'GET', 'POST', 'PUT', 'DELETE', 'PATCH', 'HEAD', 'OPTIONS'}
|
||||
if method not in allowed_methods:
|
||||
return {'response': None, 'status_code': 0, 'headers': {}, 'error': f'Invalid method: {method}'}
|
||||
|
||||
# 创建 headers 副本,避免修改输入
|
||||
headers = dict(inputs.get('headers', {}))
|
||||
headers['Content-Type'] = content_type
|
||||
|
||||
auth_type = self.get_config('auth_type', 'none')
|
||||
auth_config = self.get_config('auth_config', {})
|
||||
|
||||
if auth_type == 'bearer':
|
||||
headers['Authorization'] = f'Bearer {auth_config.get("token", "")}'
|
||||
elif auth_type == 'api_key':
|
||||
header_name = auth_config.get('header', 'X-API-Key')
|
||||
headers[header_name] = auth_config.get('key', '')
|
||||
|
||||
body = inputs.get('body')
|
||||
|
||||
logger.info('HTTP %s %s (timeout=%s)', method, url, timeout)
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.request(
|
||||
method=method,
|
||||
url=url,
|
||||
json=body if content_type == 'application/json' else None,
|
||||
data=body if content_type != 'application/json' else None,
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=timeout),
|
||||
allow_redirects=allow_redirects,
|
||||
) as response:
|
||||
try:
|
||||
response_data = await response.json()
|
||||
except Exception:
|
||||
response_data = await response.text()
|
||||
|
||||
logger.info('HTTP %s %s -> %d', method, url, response.status)
|
||||
|
||||
return {
|
||||
'response': response_data,
|
||||
'status_code': response.status,
|
||||
'headers': dict(response.headers),
|
||||
'error': None,
|
||||
}
|
||||
except aiohttp.ClientError as e:
|
||||
logger.error('HTTP request failed: %s', e)
|
||||
return {'response': None, 'status_code': 0, 'headers': {}, 'error': f'HTTP error: {e}'}
|
||||
except Exception as e:
|
||||
logger.error('HTTP request unexpected error: %s', e)
|
||||
return {'response': None, 'status_code': 0, 'headers': {}, 'error': f'Unexpected error: {e}'}
|
||||
@@ -0,0 +1,32 @@
|
||||
"""Iterator Node - Dify-style iterator for processing array items"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('iterator')
|
||||
class IteratorNode(WorkflowNode):
|
||||
"""Iterator node - iterate over array items one by one"""
|
||||
|
||||
category = 'control'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
items = inputs.get('items', [])
|
||||
if not isinstance(items, list):
|
||||
items = [items] if items else []
|
||||
|
||||
max_iterations = self.get_config('max_iterations', 1000)
|
||||
items = items[:max_iterations]
|
||||
|
||||
return {
|
||||
'item': items[0] if items else None,
|
||||
'index': 0,
|
||||
'is_first': True,
|
||||
'is_last': len(items) <= 1,
|
||||
'results': [],
|
||||
'completed': len(items) == 0,
|
||||
'_items': items,
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
"""Knowledge Retrieval Node - search in knowledge base
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/knowledge_retrieval.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('knowledge_retrieval')
|
||||
class KnowledgeRetrievalNode(WorkflowNode):
|
||||
"""Knowledge retrieval node - search in knowledge base"""
|
||||
|
||||
category = 'process'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
query = inputs.get('query', '')
|
||||
return {'documents': [], 'citations': [], 'context': f'[Knowledge base search for: {query}]'}
|
||||
@@ -0,0 +1,37 @@
|
||||
"""Langflow Flow Node - call Langflow API
|
||||
|
||||
Node metadata is loaded from: ../../templates/metadata/nodes/langflow_flow.yaml
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('langflow_flow')
|
||||
class LangflowFlowNode(WorkflowNode):
|
||||
"""Langflow flow node - call Langflow API"""
|
||||
|
||||
category = 'integration'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
base_url = self.get_config('base_url', 'http://localhost:7860')
|
||||
api_key = self.get_config('api_key', '')
|
||||
flow_id = self.get_config('flow_id', '')
|
||||
input_value = inputs.get('input_value', '')
|
||||
|
||||
# Safe API key truncation
|
||||
masked_key = f'{api_key[:4]}...{api_key[-4:]}' if len(api_key) > 8 else '***' if api_key else ''
|
||||
|
||||
return {
|
||||
'result': None,
|
||||
'success': False,
|
||||
'_debug': {
|
||||
'base_url': base_url,
|
||||
'api_key': masked_key,
|
||||
'flow_id': flow_id,
|
||||
'input_value': input_value,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,841 @@
|
||||
"""LLM Call Node - invoke large language model with Agent capabilities.
|
||||
|
||||
Supports:
|
||||
- Primary model with fallback models
|
||||
- Knowledge base retrieval with reranking
|
||||
- Max round context control
|
||||
- Streaming output
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from typing import Any, AsyncGenerator
|
||||
|
||||
import langbot_plugin.api.entities.builtin.provider.message as provider_message
|
||||
import langbot_plugin.api.entities.builtin.rag.context as rag_context
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
from .. import monitoring_helper
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Pre-compiled regex patterns for CoT content removal (performance optimization)
|
||||
_THINK_PATTERNS = [
|
||||
re.compile(r'<think>.*?</think>', re.DOTALL | re.IGNORECASE),
|
||||
re.compile(r'<thought>.*?</thought>', re.DOTALL | re.IGNORECASE),
|
||||
re.compile(r'<reasoning>.*?</reasoning>', re.DOTALL | re.IGNORECASE),
|
||||
re.compile(r'<\u601d\u8003>.*?</\u601d\u8003>', re.DOTALL | re.IGNORECASE),
|
||||
re.compile(r'<\u63a8\u7406>.*?</\u63a8\u7406>', re.DOTALL | re.IGNORECASE),
|
||||
]
|
||||
|
||||
# Template variable regex
|
||||
_TEMPLATE_VAR_RE = re.compile(r'\{\{([^}]+)\}\}')
|
||||
|
||||
|
||||
@workflow_node('llm_call')
|
||||
class LLMCallNode(WorkflowNode):
|
||||
"""LLM call node - invoke large language model"""
|
||||
|
||||
category = 'process'
|
||||
|
||||
def _resolve_template(self, template: str, inputs: dict[str, Any], context: ExecutionContext) -> str:
|
||||
"""Resolve {{variable}} placeholders in a template string."""
|
||||
if not template:
|
||||
return ''
|
||||
|
||||
unresolved_vars = []
|
||||
|
||||
def replacer(match: re.Match) -> str:
|
||||
expr = match.group(1).strip()
|
||||
# Try inputs first
|
||||
if expr in inputs:
|
||||
return str(inputs[expr])
|
||||
# Try context variables
|
||||
if expr.startswith('variables.'):
|
||||
var_name = expr[len('variables.'):]
|
||||
return str(context.variables.get(var_name, ''))
|
||||
# Try message context
|
||||
if expr.startswith('message.') and context.message_context:
|
||||
attr = expr[len('message.'):]
|
||||
return str(getattr(context.message_context, attr, ''))
|
||||
unresolved_vars.append(expr)
|
||||
return match.group(0) # leave unresolved
|
||||
|
||||
result = _TEMPLATE_VAR_RE.sub(replacer, template)
|
||||
|
||||
# Log warning for unresolved variables
|
||||
if unresolved_vars:
|
||||
logger.warning(
|
||||
f'LLM call node {self.node_id}: unresolved template variables: {unresolved_vars}'
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
def _remove_think_content(self, text: str) -> str:
|
||||
"""Remove CoT (Chain of Thought) thinking content from response."""
|
||||
if not text:
|
||||
return text
|
||||
|
||||
result = text
|
||||
for pattern in _THINK_PATTERNS:
|
||||
result = pattern.sub('', result)
|
||||
|
||||
return result.strip()
|
||||
|
||||
def _apply_content_filter(self, text: str) -> tuple[str, bool, str]:
|
||||
"""Apply content safety filter to text.
|
||||
|
||||
Returns:
|
||||
(filtered_text, is_blocked, user_notice)
|
||||
"""
|
||||
if not text or not self.ap:
|
||||
return text, False, ''
|
||||
|
||||
# Check if content filter is enabled
|
||||
safety_config = getattr(self.ap, 'pipeline_cfg', None)
|
||||
if not safety_config:
|
||||
return text, False, ''
|
||||
|
||||
# Check sensitive words
|
||||
sensitive_words = []
|
||||
try:
|
||||
if hasattr(self.ap, 'sensitive_meta') and hasattr(self.ap.sensitive_meta, 'data'):
|
||||
sensitive_words = self.ap.sensitive_meta.data.get('words', [])
|
||||
except Exception as e:
|
||||
logger.warning("Failed to load sensitive words from sensitive_meta: %s", e)
|
||||
sensitive_words = []
|
||||
|
||||
if not sensitive_words:
|
||||
return text, False, ''
|
||||
|
||||
found = False
|
||||
filtered_text = text
|
||||
for word in sensitive_words:
|
||||
try:
|
||||
matches = re.findall(word, filtered_text, re.IGNORECASE)
|
||||
if matches:
|
||||
found = True
|
||||
mask_word = ''
|
||||
mask = '*'
|
||||
try:
|
||||
if hasattr(self.ap, 'sensitive_meta') and hasattr(self.ap.sensitive_meta, 'data'):
|
||||
mask_word = self.ap.sensitive_meta.data.get('mask_word', '')
|
||||
mask = self.ap.sensitive_meta.data.get('mask', '*')
|
||||
except Exception as e:
|
||||
# Keep default mask settings when sensitive metadata is unavailable or malformed.
|
||||
logger.debug(
|
||||
f'LLM call node {self.node_id}: failed to read sensitive mask config, using defaults: {e}'
|
||||
)
|
||||
|
||||
for m in matches:
|
||||
if mask_word:
|
||||
filtered_text = filtered_text.replace(m, mask_word)
|
||||
else:
|
||||
filtered_text = filtered_text.replace(m, mask * len(m))
|
||||
except re.error:
|
||||
# Invalid regex pattern, skip
|
||||
continue
|
||||
|
||||
if found:
|
||||
return filtered_text, False, '消息中存在不合适的内容, 请修改'
|
||||
|
||||
return text, False, ''
|
||||
|
||||
# RAG combined prompt template (same as localagent.py)
|
||||
RAG_COMBINED_PROMPT_TEMPLATE = """
|
||||
The following are relevant context entries retrieved from the knowledge base.
|
||||
Please use them to answer the user's message.
|
||||
Respond in the same language as the user's input.
|
||||
|
||||
<context>
|
||||
{rag_context}
|
||||
</context>
|
||||
|
||||
<user_message>
|
||||
{user_message}
|
||||
</user_message>
|
||||
"""
|
||||
|
||||
def _build_system_prompt_with_format(self, base_prompt: str, output_format: str, json_schema: str) -> str:
|
||||
"""Build system prompt with output format instructions."""
|
||||
prompt = base_prompt
|
||||
|
||||
if output_format == 'json':
|
||||
prompt += '\n\nPlease respond in valid JSON format.'
|
||||
if json_schema:
|
||||
prompt += f'\nFollow this JSON schema:\n{json_schema}'
|
||||
elif output_format == 'markdown':
|
||||
prompt += '\n\nPlease respond in Markdown format.'
|
||||
|
||||
return prompt
|
||||
|
||||
def _build_messages_from_prompt_array(
|
||||
self,
|
||||
prompt_array: list[dict],
|
||||
inputs: dict[str, Any],
|
||||
context: ExecutionContext,
|
||||
output_format: str,
|
||||
json_schema: str,
|
||||
) -> list[provider_message.Message]:
|
||||
"""Build messages list from prompt array (same format as pipeline).
|
||||
|
||||
Each item in prompt_array is {role: str, content: str}.
|
||||
Resolves template variables in content.
|
||||
"""
|
||||
messages: list[provider_message.Message] = []
|
||||
|
||||
for item in prompt_array:
|
||||
role = item.get('role', 'user')
|
||||
content = item.get('content', '')
|
||||
|
||||
# Resolve template variables in content
|
||||
resolved_content = self._resolve_template(content, inputs, context)
|
||||
|
||||
# Apply format instructions to system prompt
|
||||
if role == 'system':
|
||||
resolved_content = self._build_system_prompt_with_format(
|
||||
resolved_content, output_format, json_schema
|
||||
)
|
||||
|
||||
messages.append(provider_message.Message(role=role, content=resolved_content))
|
||||
|
||||
return messages
|
||||
|
||||
async def _get_model_candidates(self, model_uuid: str, fallback_models: list) -> list:
|
||||
"""Build ordered list of models to try: primary model + fallback models."""
|
||||
candidates = []
|
||||
|
||||
# Primary model
|
||||
if model_uuid:
|
||||
try:
|
||||
primary = await self.ap.model_mgr.get_model_by_uuid(model_uuid)
|
||||
candidates.append(primary)
|
||||
except ValueError:
|
||||
logger.warning(f'[LLM:{self.node_id}] Primary model {model_uuid} not found')
|
||||
|
||||
# Fallback models
|
||||
for fb_uuid in fallback_models:
|
||||
try:
|
||||
fb_model = await self.ap.model_mgr.get_model_by_uuid(fb_uuid)
|
||||
candidates.append(fb_model)
|
||||
except ValueError:
|
||||
logger.warning(f'[LLM:{self.node_id}] Fallback model {fb_uuid} not found, skipping')
|
||||
|
||||
return candidates
|
||||
|
||||
async def _invoke_with_fallback(
|
||||
self,
|
||||
candidates: list,
|
||||
messages: list,
|
||||
funcs: list | None,
|
||||
extra_args: dict,
|
||||
) -> tuple[Any, Any, dict]:
|
||||
"""Try non-streaming invocation with sequential fallback. Returns (message, model_used, usage_info)."""
|
||||
last_error = None
|
||||
for model in candidates:
|
||||
try:
|
||||
result = await model.provider.invoke_llm(
|
||||
query=None,
|
||||
model=model,
|
||||
messages=messages,
|
||||
funcs=funcs if model.model_entity.abilities.__contains__('func_call') else [],
|
||||
extra_args=extra_args,
|
||||
)
|
||||
# invoke_llm returns (message, usage_info) tuple
|
||||
if isinstance(result, tuple) and len(result) == 2:
|
||||
msg, usage_info = result
|
||||
else:
|
||||
msg = result
|
||||
usage_info = {}
|
||||
return msg, model, usage_info
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
logger.warning(f'[LLM:{self.node_id}] Model {model.model_entity.name} failed: {e}, trying next...')
|
||||
raise last_error or RuntimeError('No model candidates available')
|
||||
|
||||
async def _retrieve_knowledge(
|
||||
self,
|
||||
user_message_text: str,
|
||||
knowledge_bases: list[str],
|
||||
rerank_model_uuid: str,
|
||||
rerank_top_k: int,
|
||||
) -> str:
|
||||
"""Retrieve from knowledge bases and optionally rerank results.
|
||||
|
||||
Returns the enhanced user message text with RAG context, or original text if no results.
|
||||
"""
|
||||
if not knowledge_bases or not user_message_text:
|
||||
return user_message_text
|
||||
|
||||
all_results: list[rag_context.RetrievalResultEntry] = []
|
||||
|
||||
# Retrieve from each knowledge base
|
||||
for kb_uuid in knowledge_bases:
|
||||
try:
|
||||
kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_uuid)
|
||||
if not kb:
|
||||
logger.warning(f'[LLM:{self.node_id}] Knowledge base {kb_uuid} not found, skipping')
|
||||
continue
|
||||
|
||||
result = await kb.retrieve(user_message_text, settings={})
|
||||
if result:
|
||||
all_results.extend(result)
|
||||
except Exception as e:
|
||||
logger.warning(f'[LLM:{self.node_id}] Failed to retrieve from KB {kb_uuid}: {e}')
|
||||
|
||||
# Rerank step: re-score results using a rerank model if configured
|
||||
if all_results and rerank_model_uuid:
|
||||
try:
|
||||
rerank_model = await self.ap.model_mgr.get_rerank_model_by_uuid(rerank_model_uuid)
|
||||
|
||||
doc_texts = []
|
||||
for entry in all_results:
|
||||
text = ' '.join(c.text for c in entry.content if c.type == 'text' and c.text)
|
||||
doc_texts.append(text)
|
||||
|
||||
doc_texts_capped = doc_texts[:64] # Cap for reranker input
|
||||
scores = await rerank_model.provider.invoke_rerank(
|
||||
model=rerank_model,
|
||||
query=user_message_text,
|
||||
documents=doc_texts_capped,
|
||||
)
|
||||
|
||||
scored = sorted(scores, key=lambda x: x.get('relevance_score', 0), reverse=True)
|
||||
top_indices = [s['index'] for s in scored[:rerank_top_k] if s['index'] < len(all_results)]
|
||||
all_results = [all_results[i] for i in top_indices]
|
||||
|
||||
logger.info(
|
||||
f'[LLM:{self.node_id}] Rerank complete: {len(doc_texts)} docs -> top {len(all_results)} kept (top_k={rerank_top_k})'
|
||||
)
|
||||
except ValueError:
|
||||
logger.warning(f'[LLM:{self.node_id}] Rerank model {rerank_model_uuid} not found, skipping rerank')
|
||||
except Exception as e:
|
||||
logger.warning(f'[LLM:{self.node_id}] Rerank failed, using original order: {e}')
|
||||
|
||||
# Build RAG context text
|
||||
if all_results:
|
||||
texts = []
|
||||
idx = 1
|
||||
for entry in all_results:
|
||||
for content in entry.content:
|
||||
if content.type == 'text' and content.text is not None:
|
||||
texts.append(f'[{idx}] {content.text}')
|
||||
idx += 1
|
||||
rag_context_text = '\n\n'.join(texts)
|
||||
return self.RAG_COMBINED_PROMPT_TEMPLATE.format(
|
||||
rag_context=rag_context_text,
|
||||
user_message=user_message_text,
|
||||
)
|
||||
|
||||
return user_message_text
|
||||
|
||||
def _build_messages_with_history(
|
||||
self,
|
||||
system_prompt: str,
|
||||
user_message_text: str,
|
||||
context: ExecutionContext,
|
||||
max_round: int,
|
||||
) -> list[provider_message.Message]:
|
||||
"""Build messages list with conversation history up to max_round."""
|
||||
messages: list[provider_message.Message] = []
|
||||
|
||||
# Add system prompt
|
||||
if system_prompt:
|
||||
messages.append(provider_message.Message(role='system', content=system_prompt))
|
||||
|
||||
# Get conversation history from context
|
||||
conversation_history = context.variables.get('_conversation_history', [])
|
||||
|
||||
# Apply max_round limit (each round = 1 user + 1 assistant message)
|
||||
if max_round > 0 and conversation_history:
|
||||
# Keep only the last max_round * 2 messages (user + assistant pairs)
|
||||
max_messages = max_round * 2
|
||||
if len(conversation_history) > max_messages:
|
||||
conversation_history = conversation_history[-max_messages:]
|
||||
|
||||
# Add conversation history
|
||||
for msg in conversation_history:
|
||||
if isinstance(msg, dict):
|
||||
role = msg.get('role', 'user')
|
||||
content = msg.get('content', '')
|
||||
messages.append(provider_message.Message(role=role, content=content))
|
||||
elif hasattr(msg, 'role') and hasattr(msg, 'content'):
|
||||
messages.append(provider_message.Message(role=msg.role, content=msg.content))
|
||||
|
||||
# Add current user message
|
||||
messages.append(provider_message.Message(role='user', content=user_message_text))
|
||||
|
||||
return messages
|
||||
|
||||
def _save_to_conversation_history(
|
||||
self,
|
||||
context: ExecutionContext,
|
||||
user_message_text: str,
|
||||
response_text: str,
|
||||
max_round: int,
|
||||
) -> None:
|
||||
"""Save the exchange to conversation history."""
|
||||
if max_round <= 0:
|
||||
return
|
||||
|
||||
history = context.variables.get('_conversation_history', [])
|
||||
history.append({'role': 'user', 'content': user_message_text})
|
||||
history.append({'role': 'assistant', 'content': response_text})
|
||||
|
||||
# Enforce max_round limit
|
||||
max_messages = max_round * 2
|
||||
if len(history) > max_messages:
|
||||
history = history[-max_messages:]
|
||||
|
||||
context.variables['_conversation_history'] = history
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
# Support both new model_config format and legacy model + fallback_models format
|
||||
model_config = self.get_config('model_config', None)
|
||||
if model_config and isinstance(model_config, dict):
|
||||
# New format: {primary: uuid, fallbacks: [uuid1, uuid2, ...]}
|
||||
model_uuid = model_config.get('primary', '')
|
||||
fallback_models = model_config.get('fallbacks', [])
|
||||
else:
|
||||
# Legacy format: separate model and fallback_models
|
||||
model_uuid = self.get_config('model', '')
|
||||
fallback_models = self.get_config('fallback_models', [])
|
||||
|
||||
if not model_uuid:
|
||||
raise ValueError('No model configured for LLM call node')
|
||||
|
||||
if not self.ap:
|
||||
raise RuntimeError('Application instance not available - cannot call LLM')
|
||||
|
||||
# Get error handling config
|
||||
exception_handling = self.get_config('exception_handling', 'show-error')
|
||||
failure_hint = self.get_config('failure_hint', 'Request failed.')
|
||||
track_function_calls = self.get_config('track_function_calls', False)
|
||||
|
||||
# Get output format and json_schema config
|
||||
output_format = self.get_config('output_format', 'text')
|
||||
json_schema = self.get_config('json_schema', '')
|
||||
|
||||
# Agent config: knowledge bases, rerank, max_round
|
||||
# (fallback_models already resolved above from model_config or fallback_models)
|
||||
knowledge_bases = self.get_config('knowledge_bases', [])
|
||||
rerank_model = self.get_config('rerank_model', '')
|
||||
rerank_top_k = self.get_config('rerank_top_k', 5)
|
||||
max_round = self.get_config('max_round', 10)
|
||||
|
||||
# Resolve prompts - support both new prompt array format and legacy format
|
||||
prompt_array = self.get_config('prompt')
|
||||
user_prompt = '' # Initialize for later use in _save_to_conversation_history
|
||||
|
||||
if prompt_array and isinstance(prompt_array, list):
|
||||
# New format: prompt array like pipeline
|
||||
messages = self._build_messages_from_prompt_array(
|
||||
prompt_array, inputs, context, output_format, json_schema
|
||||
)
|
||||
|
||||
# Get user input text for knowledge retrieval
|
||||
user_input = inputs.get('input', '')
|
||||
|
||||
# Knowledge retrieval: enhance user input with RAG context
|
||||
user_input = await self._retrieve_knowledge(
|
||||
user_message_text=user_input,
|
||||
knowledge_bases=knowledge_bases,
|
||||
rerank_model_uuid=rerank_model,
|
||||
rerank_top_k=rerank_top_k,
|
||||
)
|
||||
|
||||
# Track user_prompt for conversation history
|
||||
user_prompt = user_input
|
||||
|
||||
# Add user input as last message
|
||||
if user_input:
|
||||
messages.append(provider_message.Message(role='user', content=user_input))
|
||||
|
||||
# Apply max_round to conversation history
|
||||
conversation_history = context.variables.get('_conversation_history', [])
|
||||
if max_round > 0 and conversation_history:
|
||||
max_messages = max_round * 2
|
||||
if len(conversation_history) > max_messages:
|
||||
conversation_history = conversation_history[-max_messages:]
|
||||
# Insert conversation history before user input
|
||||
history_messages = []
|
||||
for msg in conversation_history:
|
||||
if isinstance(msg, dict):
|
||||
role = msg.get('role', 'user')
|
||||
content = msg.get('content', '')
|
||||
history_messages.append(provider_message.Message(role=role, content=content))
|
||||
elif hasattr(msg, 'role') and hasattr(msg, 'content'):
|
||||
history_messages.append(provider_message.Message(role=msg.role, content=msg.content))
|
||||
# Insert history before user message
|
||||
if history_messages and len(messages) > 0:
|
||||
messages = messages[:-1] + history_messages + [messages[-1]]
|
||||
else:
|
||||
# Legacy format: separate system_prompt and user_prompt_template
|
||||
system_prompt = self._resolve_template(self.get_config('system_prompt') or '', inputs, context)
|
||||
user_prompt_template = self.get_config('user_prompt_template')
|
||||
if user_prompt_template is None:
|
||||
user_prompt_template = '{{input}}'
|
||||
user_prompt = self._resolve_template(user_prompt_template, inputs, context)
|
||||
|
||||
# Build system prompt with format instructions
|
||||
system_prompt = self._build_system_prompt_with_format(system_prompt, output_format, json_schema)
|
||||
|
||||
# Knowledge retrieval: enhance user prompt with RAG context
|
||||
user_prompt = await self._retrieve_knowledge(
|
||||
user_message_text=user_prompt,
|
||||
knowledge_bases=knowledge_bases,
|
||||
rerank_model_uuid=rerank_model,
|
||||
rerank_top_k=rerank_top_k,
|
||||
)
|
||||
|
||||
# Build messages with conversation history
|
||||
messages = self._build_messages_with_history(
|
||||
system_prompt=system_prompt,
|
||||
user_message_text=user_prompt,
|
||||
context=context,
|
||||
max_round=max_round,
|
||||
)
|
||||
|
||||
# Get model candidates (primary + fallbacks)
|
||||
candidates = await self._get_model_candidates(model_uuid, fallback_models)
|
||||
if not candidates:
|
||||
raise ValueError('No valid model candidates available')
|
||||
|
||||
# Build extra args from config
|
||||
extra_args: dict[str, Any] = {}
|
||||
temperature = self.get_config('temperature')
|
||||
if temperature is not None:
|
||||
extra_args['temperature'] = float(temperature)
|
||||
max_tokens = self.get_config('max_tokens', 0)
|
||||
if max_tokens and int(max_tokens) > 0:
|
||||
extra_args['max_tokens'] = int(max_tokens)
|
||||
|
||||
# Track start time for duration calculation
|
||||
self._llm_start_time = time.time()
|
||||
|
||||
# Invoke LLM with fallback
|
||||
try:
|
||||
result_message, used_model, llm_usage = await self._invoke_with_fallback(
|
||||
candidates=candidates,
|
||||
messages=messages,
|
||||
funcs=None,
|
||||
extra_args=extra_args,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f'[LLM:{self.node_id}] LLM call failed: {e}')
|
||||
|
||||
# Handle based on exception handling strategy
|
||||
if exception_handling == 'show-error':
|
||||
raise
|
||||
elif exception_handling == 'show-hint':
|
||||
return {
|
||||
'response': failure_hint,
|
||||
'usage': {
|
||||
'prompt_tokens': 0,
|
||||
'completion_tokens': 0,
|
||||
'total_tokens': 0,
|
||||
},
|
||||
'error': str(e),
|
||||
'error_hint_shown': True,
|
||||
}
|
||||
else: # hide
|
||||
return {
|
||||
'response': '',
|
||||
'usage': {
|
||||
'prompt_tokens': 0,
|
||||
'completion_tokens': 0,
|
||||
'total_tokens': 0,
|
||||
},
|
||||
'error': str(e),
|
||||
}
|
||||
|
||||
# Extract response text
|
||||
response_text = ''
|
||||
if isinstance(result_message.content, str):
|
||||
response_text = result_message.content
|
||||
elif isinstance(result_message.content, list):
|
||||
for elem in result_message.content:
|
||||
if hasattr(elem, 'text') and elem.text:
|
||||
response_text += elem.text
|
||||
elif isinstance(elem, str):
|
||||
response_text += elem
|
||||
|
||||
# Remove CoT content (always remove to avoid leaking internal reasoning)
|
||||
response_text = self._remove_think_content(response_text)
|
||||
|
||||
# Initialize usage default
|
||||
usage = {
|
||||
'prompt_tokens': 0,
|
||||
'completion_tokens': 0,
|
||||
'total_tokens': 0,
|
||||
}
|
||||
|
||||
# Apply content safety filter
|
||||
response_text, is_blocked, filter_notice = self._apply_content_filter(response_text)
|
||||
if is_blocked:
|
||||
logger.warning(f'[LLM:{self.node_id}] Response blocked by content filter: {filter_notice}')
|
||||
return {
|
||||
'response': filter_notice,
|
||||
'usage': usage,
|
||||
'blocked_by_filter': True,
|
||||
}
|
||||
|
||||
# Extract usage info from LLM call result
|
||||
# Priority: llm_usage (from _invoke_with_fallback) > result_message.usage > result_message.token_usage
|
||||
if llm_usage:
|
||||
usage = {
|
||||
'prompt_tokens': llm_usage.get('input_tokens', 0) or llm_usage.get('prompt_tokens', 0),
|
||||
'completion_tokens': llm_usage.get('output_tokens', 0) or llm_usage.get('completion_tokens', 0),
|
||||
'total_tokens': llm_usage.get('total_tokens', 0),
|
||||
}
|
||||
# Check result_message.usage (set by RuntimeProvider.invoke_llm)
|
||||
elif hasattr(result_message, 'usage') and result_message.usage:
|
||||
u = result_message.usage
|
||||
if isinstance(u, dict):
|
||||
usage = {
|
||||
'prompt_tokens': u.get('input_tokens', 0) or u.get('prompt_tokens', 0),
|
||||
'completion_tokens': u.get('output_tokens', 0) or u.get('completion_tokens', 0),
|
||||
'total_tokens': u.get('total_tokens', 0),
|
||||
}
|
||||
else:
|
||||
usage = {
|
||||
'prompt_tokens': getattr(u, 'input_tokens', 0) or getattr(u, 'prompt_tokens', 0),
|
||||
'completion_tokens': getattr(u, 'output_tokens', 0) or getattr(u, 'completion_tokens', 0),
|
||||
'total_tokens': getattr(u, 'total_tokens', 0),
|
||||
}
|
||||
elif hasattr(result_message, 'token_usage') and result_message.token_usage:
|
||||
u = result_message.token_usage
|
||||
if isinstance(u, dict):
|
||||
usage = {
|
||||
'prompt_tokens': u.get('prompt_tokens', 0) or 0,
|
||||
'completion_tokens': u.get('completion_tokens', 0) or 0,
|
||||
'total_tokens': u.get('total_tokens', 0) or 0,
|
||||
}
|
||||
else:
|
||||
usage = {
|
||||
'prompt_tokens': getattr(u, 'prompt_tokens', 0) or 0,
|
||||
'completion_tokens': getattr(u, 'completion_tokens', 0) or 0,
|
||||
'total_tokens': getattr(u, 'total_tokens', 0) or 0,
|
||||
}
|
||||
|
||||
# Log successful response (matching Pipeline's cut_str behavior)
|
||||
def _cut_str(s: str) -> str:
|
||||
s0 = s.split('\n')[0]
|
||||
if len(s0) > 20 or '\n' in s:
|
||||
s0 = s0[:20] + '...'
|
||||
return s0
|
||||
logger.info(f'[LLM:{self.node_id}] Response: {_cut_str(response_text)}')
|
||||
|
||||
# Record LLM call log only (response log is redundant)
|
||||
try:
|
||||
if self.ap and context.query:
|
||||
workflow_id = context.workflow_id or ''
|
||||
workflow_name = context.variables.get('_workflow_name', 'Workflow')
|
||||
bot_name = context.variables.get('_bot_name', 'Workflow')
|
||||
node_name = self.get_config('name', self.node_id)
|
||||
model_name = used_model.model_entity.name if used_model else 'unknown'
|
||||
|
||||
# Calculate duration
|
||||
duration_ms = 0
|
||||
if hasattr(self, '_llm_start_time'):
|
||||
duration_ms = int((time.time() - self._llm_start_time) * 1000)
|
||||
|
||||
# Get message_id for LLM call association
|
||||
message_id = context.variables.get('_monitoring_message_id')
|
||||
|
||||
# Record LLM call log with message_id association
|
||||
await monitoring_helper.WorkflowMonitoringHelper.record_llm_call_log(
|
||||
ap=self.ap,
|
||||
query=context.query,
|
||||
workflow_id=workflow_id,
|
||||
workflow_name=workflow_name,
|
||||
node_name=node_name,
|
||||
model_name=model_name,
|
||||
input_tokens=usage.get('prompt_tokens', 0),
|
||||
output_tokens=usage.get('completion_tokens', 0),
|
||||
duration_ms=duration_ms,
|
||||
status='success',
|
||||
bot_name=bot_name,
|
||||
context_vars=context.variables,
|
||||
message_id=message_id,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f'[LLM:{self.node_id}] Failed to record LLM logs: {e}')
|
||||
|
||||
# Save to conversation history
|
||||
self._save_to_conversation_history(
|
||||
context=context,
|
||||
user_message_text=user_prompt,
|
||||
response_text=response_text,
|
||||
max_round=max_round,
|
||||
)
|
||||
|
||||
# Build result
|
||||
result: dict[str, Any] = {
|
||||
'response': response_text,
|
||||
'usage': usage,
|
||||
'model_used': used_model.model_entity.name if used_model else None,
|
||||
'model_uuid': used_model.model_entity.uuid if used_model else None,
|
||||
}
|
||||
|
||||
# Parse JSON output if format is json
|
||||
if output_format == 'json' and response_text:
|
||||
try:
|
||||
result['parsed'] = json.loads(response_text)
|
||||
except json.JSONDecodeError as e:
|
||||
logger.warning(f'[LLM:{self.node_id}] Failed to parse JSON: {e}')
|
||||
result['parsed'] = None
|
||||
result['parse_error'] = str(e)
|
||||
|
||||
# Add function call tracking info if configured
|
||||
if track_function_calls:
|
||||
result['function_calls'] = []
|
||||
|
||||
return result
|
||||
|
||||
async def execute_stream(
|
||||
self, inputs: dict[str, Any], context: ExecutionContext
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""Execute the LLM call with streaming output.
|
||||
|
||||
Yields chunks of response text as they arrive.
|
||||
Falls back to non-streaming if streaming is not available.
|
||||
"""
|
||||
# Support both new model_config format and legacy model + fallback_models format
|
||||
model_config = self.get_config('model_config', None)
|
||||
if model_config and isinstance(model_config, dict):
|
||||
model_uuid = model_config.get('primary', '')
|
||||
else:
|
||||
model_uuid = self.get_config('model', '')
|
||||
|
||||
if not model_uuid:
|
||||
raise ValueError('No model configured for LLM call node')
|
||||
|
||||
if not self.ap:
|
||||
raise RuntimeError('Application instance not available - cannot call LLM')
|
||||
|
||||
exception_handling = self.get_config('exception_handling', 'show-error')
|
||||
failure_hint = self.get_config('failure_hint', 'Request failed.')
|
||||
|
||||
# Resolve prompts - support both new prompt array format and legacy format
|
||||
prompt_array = self.get_config('prompt')
|
||||
if prompt_array and isinstance(prompt_array, list):
|
||||
# New format: prompt array like pipeline
|
||||
messages = self._build_messages_from_prompt_array(
|
||||
prompt_array, inputs, context, 'text', '' # No format instructions for streaming
|
||||
)
|
||||
|
||||
# Add user input
|
||||
user_input = inputs.get('input', '')
|
||||
if user_input:
|
||||
messages.append(provider_message.Message(role='user', content=user_input))
|
||||
else:
|
||||
# Legacy format
|
||||
system_prompt = self._resolve_template(self.get_config('system_prompt') or '', inputs, context)
|
||||
user_prompt_template = self.get_config('user_prompt_template')
|
||||
if user_prompt_template is None:
|
||||
user_prompt_template = '{{input}}'
|
||||
user_prompt = self._resolve_template(user_prompt_template, inputs, context)
|
||||
|
||||
# Build messages
|
||||
messages = []
|
||||
if system_prompt:
|
||||
messages.append(provider_message.Message(role='system', content=system_prompt))
|
||||
messages.append(provider_message.Message(role='user', content=user_prompt))
|
||||
|
||||
# Get model
|
||||
runtime_model = await self.ap.model_mgr.get_model_by_uuid(model_uuid)
|
||||
|
||||
# Build extra args
|
||||
extra_args: dict[str, Any] = {}
|
||||
temperature = self.get_config('temperature')
|
||||
if temperature is not None:
|
||||
extra_args['temperature'] = float(temperature)
|
||||
max_tokens = self.get_config('max_tokens', 0)
|
||||
if max_tokens and int(max_tokens) > 0:
|
||||
extra_args['max_tokens'] = int(max_tokens)
|
||||
|
||||
logger.info(f'[LLM:{self.node_id}] Streaming model {model_uuid}')
|
||||
|
||||
try:
|
||||
# Try streaming first
|
||||
stream = runtime_model.provider.invoke_llm_stream(
|
||||
query=None,
|
||||
model=runtime_model,
|
||||
messages=messages,
|
||||
funcs=None,
|
||||
extra_args=extra_args,
|
||||
)
|
||||
|
||||
full_response = ''
|
||||
in_think_block = False
|
||||
async for chunk in stream:
|
||||
chunk_text = ''
|
||||
if hasattr(chunk, 'content'):
|
||||
if isinstance(chunk.content, str):
|
||||
chunk_text = chunk.content
|
||||
elif isinstance(chunk.content, list):
|
||||
for elem in chunk.content:
|
||||
if hasattr(elem, 'text') and elem.text:
|
||||
chunk_text += elem.text
|
||||
elif isinstance(elem, str):
|
||||
chunk_text += elem
|
||||
|
||||
if chunk_text:
|
||||
# Filter <think> blocks in streaming mode
|
||||
if '<think>' in chunk_text or '<thought>' in chunk_text:
|
||||
in_think_block = True
|
||||
if in_think_block:
|
||||
if '</think>' in chunk_text or '</thought>' in chunk_text:
|
||||
in_think_block = False
|
||||
chunk_text = chunk_text.split('</think>')[-1].split('</thought>')[-1]
|
||||
else:
|
||||
chunk_text = ''
|
||||
|
||||
if chunk_text:
|
||||
full_response += chunk_text
|
||||
yield chunk_text
|
||||
|
||||
# Store in context for downstream nodes
|
||||
context.variables['_last_llm_response'] = full_response
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f'[LLM:{self.node_id}] Streaming failed, falling back - {e}')
|
||||
# Fallback to non-streaming
|
||||
try:
|
||||
result_message = await runtime_model.provider.invoke_llm(
|
||||
query=None,
|
||||
model=runtime_model,
|
||||
messages=messages,
|
||||
funcs=None,
|
||||
extra_args=extra_args,
|
||||
)
|
||||
response_text = self._extract_response_text(result_message)
|
||||
# Always remove <think> content in fallback
|
||||
response_text = self._remove_think_content(response_text)
|
||||
yield response_text
|
||||
context.variables['_last_llm_response'] = response_text
|
||||
except Exception as e2:
|
||||
logger.error(f'[LLM:{self.node_id}] Fallback also failed - {e2}')
|
||||
if exception_handling == 'show-hint':
|
||||
yield failure_hint
|
||||
elif exception_handling != 'hide':
|
||||
raise
|
||||
|
||||
def _extract_response_text(self, result_message: provider_message.Message) -> str:
|
||||
"""Extract response text from LLM result message."""
|
||||
response_text = ''
|
||||
if isinstance(result_message.content, str):
|
||||
response_text = result_message.content
|
||||
elif isinstance(result_message.content, list):
|
||||
for elem in result_message.content:
|
||||
if hasattr(elem, 'text') and elem.text:
|
||||
response_text += elem.text
|
||||
elif isinstance(elem, str):
|
||||
response_text += elem
|
||||
return response_text
|
||||
@@ -0,0 +1,30 @@
|
||||
"""Loop Node - iterate over items"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('loop')
|
||||
class LoopNode(WorkflowNode):
|
||||
"""Loop node - iterate over items"""
|
||||
|
||||
category = 'control'
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
items = inputs.get('items', [])
|
||||
if not isinstance(items, list):
|
||||
items = [items] if items else []
|
||||
|
||||
max_iterations = self.get_config('max_iterations', 100)
|
||||
items = items[:max_iterations]
|
||||
|
||||
return {
|
||||
'item': items[0] if items else None,
|
||||
'index': 0,
|
||||
'results': [],
|
||||
'completed': len(items) == 0,
|
||||
'_items': items,
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
"""MCP Tool Node - Invoke MCP (Model Context Protocol) tools
|
||||
|
||||
This module contains the implementation for the MCP Tool workflow node.
|
||||
Node metadata (label, description, inputs, outputs, config) is loaded from:
|
||||
../../templates/metadata/nodes/mcp_tool.yaml
|
||||
|
||||
The i18n for label and description is handled on the frontend side.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from langbot_plugin.api.entities.builtin.workflow.entities import ExecutionContext
|
||||
from ..node import WorkflowNode, workflow_node
|
||||
|
||||
@workflow_node('mcp_tool')
|
||||
class MCPToolNode(WorkflowNode):
|
||||
"""MCP tool node - invoke MCP (Model Context Protocol) tools"""
|
||||
|
||||
# Node type for registration
|
||||
|
||||
# Category and icon - these are not i18n
|
||||
category = 'integration'
|
||||
|
||||
# Name and description - i18n handled on frontend side
|
||||
# Frontend will use node type key to look up translation
|
||||
|
||||
# Inputs/outputs/config - loaded from YAML at runtime
|
||||
|
||||
async def execute(self, inputs: dict[str, Any], context: ExecutionContext) -> dict[str, Any]:
|
||||
"""Execute the MCP tool node
|
||||
|
||||
Args:
|
||||
inputs: Input data from connected nodes
|
||||
context: Execution context with workflow state
|
||||
|
||||
Returns:
|
||||
Dictionary of output values
|
||||
"""
|
||||
server_name = self.get_config('server_name', '')
|
||||
tool_name = self.get_config('tool_name', '')
|
||||
arguments_template = self.get_config('arguments_template', '')
|
||||
timeout = self.get_config('timeout', 30)
|
||||
|
||||
arguments = inputs.get('arguments', arguments_template)
|
||||
|
||||
return {
|
||||
'result': None,
|
||||
'success': False,
|
||||
'error': f"MCP tool '{server_name}/{tool_name}' not implemented yet",
|
||||
'_debug': {
|
||||
'server_name': server_name,
|
||||
'tool_name': tool_name,
|
||||
'arguments': arguments,
|
||||
'timeout': timeout,
|
||||
},
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user