ioInference.js 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431
  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 'knowledgeRetrieval': {
  104. // 优先使用手动定义的 inputs
  105. if (d.inputs && d.inputs.length) return d.inputs
  106. // 回退:从检索语句模板提取
  107. const krVars = new Set(extractTemplateVariables(d.query))
  108. return [...krVars].map(name => ({ name, label: name, type: 'string', description: '' }))
  109. }
  110. case 'output':
  111. return d.fields || []
  112. case 'condition': {
  113. // 优先使用手动定义的 inputs
  114. if (d.inputs && d.inputs.length) return d.inputs
  115. // 回退:从条件表达式提取
  116. const vars = new Set()
  117. for (const c of (d.conditions || [])) {
  118. for (const v of extractTemplateVariables(c.expression)) {
  119. vars.add(v)
  120. }
  121. }
  122. return [...vars].map(name => ({ name, label: name, type: 'string', description: '' }))
  123. }
  124. default:
  125. return d.inputs || []
  126. }
  127. }
  128. // ========== 类型兼容性 ==========
  129. /**
  130. * 判断源类型是否可以赋值给目标类型
  131. */
  132. export function isTypeCompatible(sourceType, targetType) {
  133. if (sourceType === targetType) return true
  134. // 以下类型可隐式转为 string
  135. const stringLike = ['filePath', 'directoryPath', 'number', 'boolean']
  136. if (stringLike.includes(sourceType) && targetType === 'string') return true
  137. // array → object 兼容
  138. if (sourceType === 'array' && targetType === 'object') return true
  139. return false
  140. }
  141. // ========== 单边匹配 ==========
  142. /**
  143. * 匹配源节点输出与目标节点输入
  144. * @returns {{ status, mapping, unmatchedSource, unmatchedTarget, typeMismatches }}
  145. */
  146. export function matchIO(sourceOutputs, targetInputs, existingMapping) {
  147. const mapping = []
  148. const unmatchedSource = [...sourceOutputs]
  149. const unmatchedTarget = [...targetInputs]
  150. const typeMismatches = []
  151. // 第一轮:按变量名精确匹配
  152. for (let si = unmatchedSource.length - 1; si >= 0; si--) {
  153. const srcField = unmatchedSource[si]
  154. const ti = unmatchedTarget.findIndex(t => t.name === srcField.name)
  155. if (ti !== -1) {
  156. const tgtField = unmatchedTarget[ti]
  157. if (isTypeCompatible(srcField.type, tgtField.type)) {
  158. mapping.push({ sourceField: srcField.name, targetField: tgtField.name })
  159. } else {
  160. typeMismatches.push({ source: srcField, target: tgtField })
  161. }
  162. unmatchedSource.splice(si, 1)
  163. unmatchedTarget.splice(ti, 1)
  164. }
  165. }
  166. // 第二轮:保留已有的手动映射(不与新映射冲突的部分)
  167. if (existingMapping) {
  168. for (const em of existingMapping) {
  169. if (mapping.some(m => m.sourceField === em.sourceField && m.targetField === em.targetField)) continue
  170. if (mapping.some(m => m.sourceField === em.sourceField || m.targetField === em.targetField)) continue
  171. mapping.push(em)
  172. const ti = unmatchedTarget.findIndex(t => t.name === em.targetField)
  173. if (ti !== -1) unmatchedTarget.splice(ti, 1)
  174. }
  175. }
  176. // 判定状态
  177. let status
  178. if (targetInputs.length === 0 && sourceOutputs.length === 0) {
  179. status = 'unmapped'
  180. } else if (unmatchedTarget.length === 0 && typeMismatches.length === 0) {
  181. status = 'ok'
  182. } else if (typeMismatches.length > 0) {
  183. status = 'mismatch'
  184. } else if (unmatchedTarget.length > 0 && unmatchedSource.length === 0) {
  185. status = 'partial'
  186. } else if (unmatchedTarget.length === 0 && unmatchedSource.length > 0) {
  187. status = 'ok' // 源有多余输出,但目标全部满足
  188. } else {
  189. status = 'mismatch'
  190. }
  191. return { status, mapping, unmatchedSource, unmatchedTarget, typeMismatches }
  192. }
  193. // ========== 推断调度 ==========
  194. /**
  195. * 推断一条边的映射状态
  196. * mapping 仍然基于直接源节点的输出(精确映射),
  197. * 但状态判断基于所有可达前驱节点的合并输出(因为上下文是累积的)
  198. * @param {object} sourceNode - 源节点
  199. * @param {object} targetNode - 目标节点
  200. * @param {object} [existingEdge] - 已有边数据
  201. * @param {Array} [allNodes] - 全部节点(用于可达前驱计算)
  202. * @param {Array} [allEdges] - 全部边(用于可达前驱计算)
  203. * @returns {{ status, mapping, unmatchedSource, unmatchedTarget, typeMismatches }}
  204. */
  205. export function inferEdgeMapping(sourceNode, targetNode, existingEdge, allNodes, allEdges) {
  206. const sourceOutputs = getNodeOutputs(sourceNode)
  207. const targetInputs = getNodeInputs(targetNode)
  208. const existingMapping = existingEdge?.data?.mapping
  209. // 直接源→目标的精确映射(用于 mapping 字段)
  210. const directResult = matchIO(sourceOutputs, targetInputs, existingMapping)
  211. // 如果没有传入全图数据,降级为只看直接源
  212. if (!allNodes || !allEdges) {
  213. return directResult
  214. }
  215. // 收集所有可达前驱节点的合并输出(用于状态判断)
  216. const reachableOutputs = collectReachableOutputs(targetNode.id, allNodes, allEdges)
  217. if (reachableOutputs.length === 0 && targetInputs.length === 0) {
  218. return { ...directResult, status: 'unmapped' }
  219. }
  220. // 用合并输出判断目标输入是否全部满足
  221. const unmatched = []
  222. const typeMismatches = []
  223. for (const tgtField of targetInputs) {
  224. const srcField = reachableOutputs.find(o => o.name === tgtField.name)
  225. if (!srcField) {
  226. unmatched.push(tgtField)
  227. } else if (!isTypeCompatible(srcField.type, tgtField.type)) {
  228. typeMismatches.push({ source: srcField, target: tgtField })
  229. }
  230. }
  231. let status
  232. if (targetInputs.length === 0 && sourceOutputs.length === 0) {
  233. status = 'unmapped'
  234. } else if (unmatched.length === 0 && typeMismatches.length === 0) {
  235. status = 'ok'
  236. } else if (typeMismatches.length > 0) {
  237. status = 'mismatch'
  238. } else {
  239. status = 'partial'
  240. }
  241. return { ...directResult, status }
  242. }
  243. /**
  244. * 收集目标节点所有可达前驱节点的合并输出(BFS 回溯边图)
  245. * @param {string} targetId - 目标节点 ID
  246. * @param {Array} allNodes - 全部节点
  247. * @param {Array} allEdges - 全部边
  248. * @returns {Array} 合并后的输出字段列表(去重,先出现的优先)
  249. */
  250. function collectReachableOutputs(targetId, allNodes, allEdges) {
  251. const nodeMap = new Map(allNodes.map(n => [n.id, n]))
  252. const visited = new Set()
  253. const queue = [targetId]
  254. const merged = new Map() // name → field
  255. while (queue.length > 0) {
  256. const currentId = queue.shift()
  257. if (visited.has(currentId)) continue
  258. visited.add(currentId)
  259. // 找到所有指向当前节点的边
  260. for (const edge of allEdges) {
  261. if (edge.target === currentId && !visited.has(edge.source)) {
  262. const sourceNode = nodeMap.get(edge.source)
  263. if (sourceNode) {
  264. for (const field of getNodeOutputs(sourceNode)) {
  265. if (!merged.has(field.name)) {
  266. merged.set(field.name, field)
  267. }
  268. }
  269. queue.push(edge.source)
  270. }
  271. }
  272. }
  273. }
  274. return [...merged.values()]
  275. }
  276. /**
  277. * 推断指定边的样式
  278. */
  279. export function getEdgeStyle(status) {
  280. return EDGE_STATUS_STYLE[status] || EDGE_STATUS_STYLE.unmapped
  281. }
  282. /**
  283. * 刷新图中与指定节点关联的所有边的映射
  284. * 返回需要更新的边列表
  285. * @param {string} nodeId - 变更的节点 ID
  286. * @param {Array} nodes - 所有节点
  287. * @param {Array} edges - 所有边
  288. * @returns {Array} - 需要更新的边 [{ id, data, style }]
  289. */
  290. export function refreshMappingsForNode(nodeId, nodes, edges) {
  291. const updates = []
  292. const nodeMap = new Map(nodes.map(n => [n.id, n]))
  293. for (const edge of edges) {
  294. const isRelevant = edge.source === nodeId || edge.target === nodeId
  295. if (!isRelevant) continue
  296. const sourceNode = nodeMap.get(edge.source)
  297. const targetNode = nodeMap.get(edge.target)
  298. if (!sourceNode || !targetNode) continue
  299. const result = inferEdgeMapping(sourceNode, targetNode, edge, nodes, edges)
  300. updates.push({
  301. id: edge.id,
  302. data: {
  303. ...edge.data,
  304. mapping: result.mapping,
  305. status: result.status,
  306. unmatchedSource: result.unmatchedSource,
  307. unmatchedTarget: result.unmatchedTarget,
  308. typeMismatches: result.typeMismatches
  309. },
  310. style: getEdgeStyle(result.status)
  311. })
  312. }
  313. return updates
  314. }
  315. /**
  316. * 推断一条新边的映射(用于 onConnect)
  317. */
  318. export function inferNewEdge(sourceNode, targetNode, allNodes, allEdges) {
  319. const result = inferEdgeMapping(sourceNode, targetNode, undefined, allNodes, allEdges)
  320. return {
  321. data: {
  322. mapping: result.mapping,
  323. status: result.status,
  324. unmatchedSource: result.unmatchedSource,
  325. unmatchedTarget: result.unmatchedTarget,
  326. typeMismatches: result.typeMismatches
  327. },
  328. style: getEdgeStyle(result.status)
  329. }
  330. }
  331. // ========== 多源汇聚分配 ==========
  332. /**
  333. * 计算多源汇聚时的变量分配方案
  334. * @param {Array} sources - 源节点列表
  335. * @param {object} target - 目标节点
  336. * @param {Array} edges - 连接到目标的边列表
  337. * @returns {{ edges: Array<{ edgeId, mapping, provided, missing }> }}
  338. */
  339. export function resolveMultiSource(sources, target, edges) {
  340. const targetInputs = getNodeInputs(target)
  341. if (targetInputs.length === 0) {
  342. return { edges: edges.map(e => ({ edgeId: e.id, mapping: [], provided: [], missing: [] })) }
  343. }
  344. const allSourceOutputs = sources.map(s => getNodeOutputs(s))
  345. const results = []
  346. // 计算所有源的合并输出
  347. const mergedOutputs = new Map()
  348. for (let i = 0; i < sources.length; i++) {
  349. for (const field of allSourceOutputs[i]) {
  350. if (!mergedOutputs.has(field.name)) {
  351. mergedOutputs.set(field.name, { field, sourceIndex: i })
  352. }
  353. }
  354. }
  355. // 尝试匹配每个目标输入字段
  356. const assigned = new Map() // targetField -> sourceIndex
  357. for (const tgtField of targetInputs) {
  358. const entry = mergedOutputs.get(tgtField.name)
  359. if (entry && isTypeCompatible(entry.field.type, tgtField.type)) {
  360. assigned.set(tgtField.name, entry.sourceIndex)
  361. }
  362. }
  363. // 分配回每条边
  364. for (let i = 0; i < edges.length; i++) {
  365. const sourceOutputs = allSourceOutputs[i] || []
  366. const edgeMapping = []
  367. const provided = []
  368. for (const tgtField of targetInputs) {
  369. if (assigned.get(tgtField.name) === i) {
  370. const srcField = sourceOutputs.find(f => f.name === tgtField.name)
  371. if (srcField) {
  372. edgeMapping.push({ sourceField: srcField.name, targetField: tgtField.name })
  373. provided.push(tgtField.name)
  374. }
  375. }
  376. }
  377. const missing = targetInputs.filter(t => !assigned.has(t.name)).map(t => t.name)
  378. results.push({
  379. edgeId: edges[i].id,
  380. mapping: edgeMapping,
  381. provided,
  382. missing
  383. })
  384. }
  385. return { edges: results }
  386. }