ioInference.js 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505
  1. /**
  2. * 工作流节点 IO 变量推断引擎
  3. *
  4. * 职责:
  5. * 1. 从节点数据中提取 inputs / outputs 定义
  6. * 2. 连线时自动匹配源节点输出 → 目标节点输入
  7. * 3. 多源汇聚时计算变量分配
  8. */
  9. // ========== 类型常量 ==========
  10. export const IO_TYPES = [
  11. { label: '字符串', value: 'string' },
  12. { label: '数字', value: 'number' },
  13. { label: '布尔值', value: 'boolean' },
  14. { label: '数组', value: 'array' },
  15. { label: '对象', value: 'object' },
  16. { label: '文件路径', value: 'filePath' },
  17. { label: '目录路径', value: 'directoryPath' }
  18. ]
  19. // 边状态 → 样式
  20. export const EDGE_STATUS_STYLE = {
  21. ok: { stroke: '#22c55e', strokeWidth: 2 },
  22. unmapped: { stroke: '#666', strokeWidth: 2 },
  23. partial: { stroke: '#f59e0b', strokeWidth: 2 },
  24. mismatch: { stroke: '#ef4444', strokeWidth: 2 }
  25. }
  26. // ========== 模板变量提取 ==========
  27. /**
  28. * 从 {{变量名}} 模板中提取变量名列表(去重)
  29. */
  30. export function extractTemplateVariables(text) {
  31. if (!text) return []
  32. const matches = text.matchAll(/\{\{\s*([\w.]+)\s*\}\}/g)
  33. const seen = new Set()
  34. const result = []
  35. for (const m of matches) {
  36. if (!seen.has(m[1])) {
  37. seen.add(m[1])
  38. result.push(m[1])
  39. }
  40. }
  41. return result
  42. }
  43. // ========== 节点 IO 获取 ==========
  44. /**
  45. * 获取节点的输出字段列表
  46. */
  47. export function getNodeOutputs(node) {
  48. const d = node.data || {}
  49. switch (node.type) {
  50. case 'userInput':
  51. return (d.variables || []).map(v => ({
  52. name: v.name,
  53. label: v.label || v.name,
  54. type: v.type || 'string',
  55. description: v.description || ''
  56. }))
  57. case 'llm':
  58. if (d.outputs && d.outputs.length) return d.outputs
  59. return [{ name: 'result', label: 'LLM 输出', type: 'string', description: '大模型响应文本' }]
  60. case 'agent':
  61. return d.outputs || []
  62. case 'skill':
  63. return d.outputs || []
  64. case 'knowledgeRetrieval':
  65. if (d.outputs && d.outputs.length) return d.outputs
  66. return [
  67. { name: 'evidences', label: '检索证据列表', type: 'array', description: '命中的知识片段集合' },
  68. { name: 'evidenceCount', label: '证据数量', type: 'number', description: '命中证据条数' },
  69. { name: 'sourceType', label: '来源类型', type: 'string', description: 'DOCUMENT / STRUCTURED_DATA / GRAPH / HYBRID' },
  70. { name: 'diagnostics', label: '诊断信息', type: 'object', description: '检索过程的诊断信息' }
  71. ]
  72. case 'output':
  73. return []
  74. case 'condition':
  75. // 条件节点透传:输出 = 所有输入
  76. return getNodeInputs(node)
  77. default:
  78. return d.outputs || []
  79. }
  80. }
  81. /**
  82. * 获取节点的输入字段列表
  83. */
  84. export function getNodeInputs(node) {
  85. const d = node.data || {}
  86. switch (node.type) {
  87. case 'userInput':
  88. return []
  89. case 'llm': {
  90. // 优先使用手动定义的 inputs
  91. if (d.inputs && d.inputs.length) return d.inputs
  92. // 回退:从模板提取
  93. const vars = new Set([
  94. ...extractTemplateVariables(d.systemPrompt),
  95. ...extractTemplateVariables(d.userPrompt)
  96. ])
  97. return [...vars].map(name => ({ name, label: name, type: 'string', description: '' }))
  98. }
  99. case 'agent':
  100. return d.inputs || []
  101. case 'skill':
  102. return d.inputs || []
  103. case 'smartAction': {
  104. // 优先使用手动定义的 inputs
  105. if (d.inputs && d.inputs.length) return d.inputs
  106. // 回退:从操作要求模板提取
  107. const saVars = new Set(extractTemplateVariables(d.actionPrompt))
  108. return [...saVars].map(name => ({ name, label: name, type: 'string', description: '' }))
  109. }
  110. case 'knowledgeRetrieval': {
  111. // 优先使用手动定义的 inputs
  112. if (d.inputs && d.inputs.length) return d.inputs
  113. // 回退:从检索语句模板提取
  114. const krVars = new Set(extractTemplateVariables(d.query))
  115. return [...krVars].map(name => ({ name, label: name, type: 'string', description: '' }))
  116. }
  117. case 'output':
  118. return d.fields || []
  119. case 'condition': {
  120. // 优先使用手动定义的 inputs
  121. if (d.inputs && d.inputs.length) return d.inputs
  122. // 回退:从条件表达式提取
  123. const vars = new Set()
  124. for (const c of (d.conditions || [])) {
  125. for (const v of extractTemplateVariables(c.expression)) {
  126. vars.add(v)
  127. }
  128. }
  129. return [...vars].map(name => ({ name, label: name, type: 'string', description: '' }))
  130. }
  131. default:
  132. return d.inputs || []
  133. }
  134. }
  135. // ========== 类型兼容性 ==========
  136. /**
  137. * 判断源类型是否可以赋值给目标类型
  138. */
  139. export function isTypeCompatible(sourceType, targetType) {
  140. if (sourceType === targetType) return true
  141. // 以下类型可隐式转为 string
  142. const stringLike = ['filePath', 'directoryPath', 'number', 'boolean']
  143. if (stringLike.includes(sourceType) && targetType === 'string') return true
  144. // array → object 兼容
  145. if (sourceType === 'array' && targetType === 'object') return true
  146. return false
  147. }
  148. // ========== 单边匹配 ==========
  149. /**
  150. * 匹配源节点输出与目标节点输入
  151. * @returns {{ status, mapping, unmatchedSource, unmatchedTarget, typeMismatches }}
  152. */
  153. export function matchIO(sourceOutputs, targetInputs, existingMapping) {
  154. const mapping = []
  155. const unmatchedSource = [...sourceOutputs]
  156. const unmatchedTarget = [...targetInputs]
  157. const typeMismatches = []
  158. // 第一轮:按变量名精确匹配
  159. for (let si = unmatchedSource.length - 1; si >= 0; si--) {
  160. const srcField = unmatchedSource[si]
  161. const ti = unmatchedTarget.findIndex(t => t.name === srcField.name)
  162. if (ti !== -1) {
  163. const tgtField = unmatchedTarget[ti]
  164. if (isTypeCompatible(srcField.type, tgtField.type)) {
  165. mapping.push({ sourceField: srcField.name, targetField: tgtField.name })
  166. } else {
  167. typeMismatches.push({ source: srcField, target: tgtField })
  168. }
  169. unmatchedSource.splice(si, 1)
  170. unmatchedTarget.splice(ti, 1)
  171. }
  172. }
  173. // 第二轮:保留已有的手动映射(不与新映射冲突的部分)
  174. if (existingMapping) {
  175. for (const em of existingMapping) {
  176. if (mapping.some(m => m.sourceField === em.sourceField && m.targetField === em.targetField)) continue
  177. if (mapping.some(m => m.sourceField === em.sourceField || m.targetField === em.targetField)) continue
  178. mapping.push(em)
  179. const ti = unmatchedTarget.findIndex(t => t.name === em.targetField)
  180. if (ti !== -1) unmatchedTarget.splice(ti, 1)
  181. }
  182. }
  183. // 判定状态
  184. let status
  185. if (targetInputs.length === 0 && sourceOutputs.length === 0) {
  186. status = 'unmapped'
  187. } else if (unmatchedTarget.length === 0 && typeMismatches.length === 0) {
  188. status = 'ok'
  189. } else if (typeMismatches.length > 0) {
  190. status = 'mismatch'
  191. } else if (unmatchedTarget.length > 0 && unmatchedSource.length === 0) {
  192. status = 'partial'
  193. } else if (unmatchedTarget.length === 0 && unmatchedSource.length > 0) {
  194. status = 'ok' // 源有多余输出,但目标全部满足
  195. } else {
  196. status = 'mismatch'
  197. }
  198. return { status, mapping, unmatchedSource, unmatchedTarget, typeMismatches }
  199. }
  200. // ========== 推断调度 ==========
  201. /**
  202. * 推断一条边的映射状态
  203. * mapping 仍然基于直接源节点的输出(精确映射),
  204. * 但状态判断基于所有可达前驱节点的合并输出(因为上下文是累积的)
  205. * @param {object} sourceNode - 源节点
  206. * @param {object} targetNode - 目标节点
  207. * @param {object} [existingEdge] - 已有边数据
  208. * @param {Array} [allNodes] - 全部节点(用于可达前驱计算)
  209. * @param {Array} [allEdges] - 全部边(用于可达前驱计算)
  210. * @returns {{ status, mapping, unmatchedSource, unmatchedTarget, typeMismatches }}
  211. */
  212. export function inferEdgeMapping(sourceNode, targetNode, existingEdge, allNodes, allEdges) {
  213. const sourceOutputs = getNodeOutputs(sourceNode)
  214. const targetInputs = getNodeInputs(targetNode)
  215. const existingMapping = existingEdge?.data?.mapping
  216. // 直接源→目标的精确映射(用于 mapping 字段)
  217. const directResult = matchIO(sourceOutputs, targetInputs, existingMapping)
  218. // 如果没有传入全图数据,降级为只看直接源
  219. if (!allNodes || !allEdges) {
  220. return directResult
  221. }
  222. // 收集所有可达前驱节点的合并输出(用于状态判断)
  223. const reachableOutputs = collectReachableOutputs(targetNode.id, allNodes, allEdges)
  224. if (reachableOutputs.length === 0 && targetInputs.length === 0) {
  225. return { ...directResult, status: 'unmapped' }
  226. }
  227. // 用合并输出判断目标输入是否全部满足
  228. const unmatched = []
  229. const typeMismatches = []
  230. for (const tgtField of targetInputs) {
  231. const srcField = reachableOutputs.find(o => o.name === tgtField.name)
  232. if (!srcField) {
  233. unmatched.push(tgtField)
  234. } else if (!isTypeCompatible(srcField.type, tgtField.type)) {
  235. typeMismatches.push({ source: srcField, target: tgtField })
  236. }
  237. }
  238. let status
  239. if (targetInputs.length === 0 && sourceOutputs.length === 0) {
  240. status = 'unmapped'
  241. } else if (unmatched.length === 0 && typeMismatches.length === 0) {
  242. status = 'ok'
  243. } else if (typeMismatches.length > 0) {
  244. status = 'mismatch'
  245. } else {
  246. status = 'partial'
  247. }
  248. return { ...directResult, status }
  249. }
  250. /**
  251. * 收集目标节点所有可达前驱节点的合并输出(BFS 回溯边图)
  252. * @param {string} targetId - 目标节点 ID
  253. * @param {Array} allNodes - 全部节点
  254. * @param {Array} allEdges - 全部边
  255. * @returns {Array} 合并后的输出字段列表(去重,先出现的优先)
  256. */
  257. function collectReachableOutputs(targetId, allNodes, allEdges) {
  258. const nodeMap = new Map(allNodes.map(n => [n.id, n]))
  259. const visited = new Set()
  260. const queue = [targetId]
  261. const merged = new Map() // name → field
  262. while (queue.length > 0) {
  263. const currentId = queue.shift()
  264. if (visited.has(currentId)) continue
  265. visited.add(currentId)
  266. // 找到所有指向当前节点的边
  267. for (const edge of allEdges) {
  268. if (edge.target === currentId && !visited.has(edge.source)) {
  269. const sourceNode = nodeMap.get(edge.source)
  270. if (sourceNode) {
  271. for (const field of getNodeOutputs(sourceNode)) {
  272. if (!merged.has(field.name)) {
  273. merged.set(field.name, field)
  274. }
  275. }
  276. queue.push(edge.source)
  277. }
  278. }
  279. }
  280. }
  281. return [...merged.values()]
  282. }
  283. /**
  284. * 获取目标节点的所有可达前驱节点(通过边反向 BFS)
  285. *
  286. * 用于"前置数据关联"下拉选项:节点输入不仅可关联直接上游,
  287. * 也可关联画布中所有可达的前序节点输出。
  288. *
  289. * @param {string} targetId - 目标节点 ID
  290. * @param {Array} allNodes - 全部节点
  291. * @param {Array} allEdges - 全部边
  292. * @returns {Array} 可达前驱节点列表(按 BFS 发现顺序,不含 targetId 本身,不重复)
  293. */
  294. export function getReachablePredecessors(targetId, allNodes, allEdges) {
  295. const nodeMap = new Map(allNodes.map(n => [n.id, n]))
  296. const visited = new Set([targetId])
  297. const queue = [targetId]
  298. const result = []
  299. while (queue.length > 0) {
  300. const currentId = queue.shift()
  301. for (const edge of allEdges) {
  302. if (edge.target !== currentId) continue
  303. if (visited.has(edge.source)) continue
  304. visited.add(edge.source)
  305. const sourceNode = nodeMap.get(edge.source)
  306. if (sourceNode) {
  307. result.push(sourceNode)
  308. queue.push(edge.source)
  309. }
  310. }
  311. }
  312. return result
  313. }
  314. /**
  315. * 获取起始节点的所有可达后继节点(通过边正向 BFS)
  316. *
  317. * 用于"输出变量引入":节点输出可关联到所有可达后继节点的输入字段。
  318. *
  319. * @param {string} sourceId - 起始节点 ID
  320. * @param {Array} allNodes - 全部节点
  321. * @param {Array} allEdges - 全部边
  322. * @returns {Array} 可达后继节点列表(按 BFS 发现顺序,不含 sourceId 本身,不重复)
  323. */
  324. export function getReachableSuccessors(sourceId, allNodes, allEdges) {
  325. const nodeMap = new Map(allNodes.map(n => [n.id, n]))
  326. const visited = new Set([sourceId])
  327. const queue = [sourceId]
  328. const result = []
  329. while (queue.length > 0) {
  330. const currentId = queue.shift()
  331. for (const edge of allEdges) {
  332. if (edge.source !== currentId) continue
  333. if (visited.has(edge.target)) continue
  334. visited.add(edge.target)
  335. const targetNode = nodeMap.get(edge.target)
  336. if (targetNode) {
  337. result.push(targetNode)
  338. queue.push(edge.target)
  339. }
  340. }
  341. }
  342. return result
  343. }
  344. /**
  345. * 推断指定边的样式
  346. */
  347. export function getEdgeStyle(status) {
  348. return EDGE_STATUS_STYLE[status] || EDGE_STATUS_STYLE.unmapped
  349. }
  350. /**
  351. * 刷新图中与指定节点关联的所有边的映射
  352. * 返回需要更新的边列表
  353. * @param {string} nodeId - 变更的节点 ID
  354. * @param {Array} nodes - 所有节点
  355. * @param {Array} edges - 所有边
  356. * @returns {Array} - 需要更新的边 [{ id, data, style }]
  357. */
  358. export function refreshMappingsForNode(nodeId, nodes, edges) {
  359. const updates = []
  360. const nodeMap = new Map(nodes.map(n => [n.id, n]))
  361. for (const edge of edges) {
  362. const isRelevant = edge.source === nodeId || edge.target === nodeId
  363. if (!isRelevant) continue
  364. const sourceNode = nodeMap.get(edge.source)
  365. const targetNode = nodeMap.get(edge.target)
  366. if (!sourceNode || !targetNode) continue
  367. const result = inferEdgeMapping(sourceNode, targetNode, edge, nodes, edges)
  368. updates.push({
  369. id: edge.id,
  370. data: {
  371. ...edge.data,
  372. mapping: result.mapping,
  373. status: result.status,
  374. unmatchedSource: result.unmatchedSource,
  375. unmatchedTarget: result.unmatchedTarget,
  376. typeMismatches: result.typeMismatches
  377. },
  378. style: getEdgeStyle(result.status)
  379. })
  380. }
  381. return updates
  382. }
  383. /**
  384. * 推断一条新边的映射(用于 onConnect)
  385. */
  386. export function inferNewEdge(sourceNode, targetNode, allNodes, allEdges) {
  387. const result = inferEdgeMapping(sourceNode, targetNode, undefined, allNodes, allEdges)
  388. return {
  389. data: {
  390. mapping: result.mapping,
  391. status: result.status,
  392. unmatchedSource: result.unmatchedSource,
  393. unmatchedTarget: result.unmatchedTarget,
  394. typeMismatches: result.typeMismatches
  395. },
  396. style: getEdgeStyle(result.status)
  397. }
  398. }
  399. // ========== 多源汇聚分配 ==========
  400. /**
  401. * 计算多源汇聚时的变量分配方案
  402. * @param {Array} sources - 源节点列表
  403. * @param {object} target - 目标节点
  404. * @param {Array} edges - 连接到目标的边列表
  405. * @returns {{ edges: Array<{ edgeId, mapping, provided, missing }> }}
  406. */
  407. export function resolveMultiSource(sources, target, edges) {
  408. const targetInputs = getNodeInputs(target)
  409. if (targetInputs.length === 0) {
  410. return { edges: edges.map(e => ({ edgeId: e.id, mapping: [], provided: [], missing: [] })) }
  411. }
  412. const allSourceOutputs = sources.map(s => getNodeOutputs(s))
  413. const results = []
  414. // 计算所有源的合并输出
  415. const mergedOutputs = new Map()
  416. for (let i = 0; i < sources.length; i++) {
  417. for (const field of allSourceOutputs[i]) {
  418. if (!mergedOutputs.has(field.name)) {
  419. mergedOutputs.set(field.name, { field, sourceIndex: i })
  420. }
  421. }
  422. }
  423. // 尝试匹配每个目标输入字段
  424. const assigned = new Map() // targetField -> sourceIndex
  425. for (const tgtField of targetInputs) {
  426. const entry = mergedOutputs.get(tgtField.name)
  427. if (entry && isTypeCompatible(entry.field.type, tgtField.type)) {
  428. assigned.set(tgtField.name, entry.sourceIndex)
  429. }
  430. }
  431. // 分配回每条边
  432. for (let i = 0; i < edges.length; i++) {
  433. const sourceOutputs = allSourceOutputs[i] || []
  434. const edgeMapping = []
  435. const provided = []
  436. for (const tgtField of targetInputs) {
  437. if (assigned.get(tgtField.name) === i) {
  438. const srcField = sourceOutputs.find(f => f.name === tgtField.name)
  439. if (srcField) {
  440. edgeMapping.push({ sourceField: srcField.name, targetField: tgtField.name })
  441. provided.push(tgtField.name)
  442. }
  443. }
  444. }
  445. const missing = targetInputs.filter(t => !assigned.has(t.name)).map(t => t.name)
  446. results.push({
  447. edgeId: edges[i].id,
  448. mapping: edgeMapping,
  449. provided,
  450. missing
  451. })
  452. }
  453. return { edges: results }
  454. }