unsloth/studio/frontend/src/features/canvas-lab/utils/graph.ts
2026-02-05 21:55:50 +01:00

159 lines
4.9 KiB
TypeScript

import { type Connection, type Edge, addEdge } from "@xyflow/react";
import type { NodeConfig, SamplerConfig } from "../types";
import { HANDLE_IDS } from "./handles";
import {
isCategoryConfig,
isExpressionConfig,
isLlmConfig,
isSubcategoryConfig,
} from "./index";
function buildTemplateWithRef(template: string, ref: string): string {
if (template.includes(ref)) {
return template;
}
if (template.trim()) {
return `${template}\n${ref}`;
}
return ref;
}
function syncSubcategoryMapping(
subcategory: SamplerConfig,
parent: NodeConfig,
): SamplerConfig {
if (!isCategoryConfig(parent)) {
return {
...subcategory,
// biome-ignore lint/style/useNamingConvention: api schema
subcategory_parent: parent.name,
};
}
const nextMapping: Record<string, string[]> = {
...(subcategory.subcategory_mapping ?? {}),
};
for (const value of parent.values ?? []) {
if (!nextMapping[value]) {
nextMapping[value] = [];
}
}
return {
...subcategory,
// biome-ignore lint/style/useNamingConvention: api schema
subcategory_parent: parent.name,
// biome-ignore lint/style/useNamingConvention: api schema
subcategory_mapping: nextMapping,
};
}
function isSemanticRelation(source: NodeConfig, target: NodeConfig): boolean {
if (source.kind === "model_provider" && target.kind === "model_config") {
return true;
}
return source.kind === "model_config" && target.kind === "llm";
}
function isModelInfraNode(config: NodeConfig): boolean {
return config.kind === "model_provider" || config.kind === "model_config";
}
function isSemanticLane(connection: Connection): boolean {
return (
connection.sourceHandle === HANDLE_IDS.semanticOut &&
connection.targetHandle === HANDLE_IDS.semanticIn
);
}
function isDataLane(connection: Connection): boolean {
return (
connection.sourceHandle === HANDLE_IDS.dataOut &&
connection.targetHandle === HANDLE_IDS.dataIn
);
}
export function isValidCanvasConnection(
connection: Connection,
configs: Record<string, NodeConfig>,
): boolean {
if (!(connection.source && connection.target)) {
return false;
}
if (connection.source === connection.target) {
return false;
}
const source = configs[connection.source];
const target = configs[connection.target];
if (!(source && target)) {
return false;
}
const semanticRelation = isSemanticRelation(source, target);
if (semanticRelation) {
return isSemanticLane(connection);
}
if (isModelInfraNode(source) || isModelInfraNode(target)) {
return false;
}
return isDataLane(connection);
}
export function applyCanvasConnection(
connection: Connection,
configs: Record<string, NodeConfig>,
edges: Edge[],
): { edges: Edge[]; configs?: Record<string, NodeConfig> } {
if (!isValidCanvasConnection(connection, configs)) {
return { edges };
}
const source = connection.source ? configs[connection.source] : null;
const target = connection.target ? configs[connection.target] : null;
if (!(source && target)) {
return { edges };
}
const semanticRelation = isSemanticRelation(source, target);
const nextEdges = addEdge(
{ ...connection, type: semanticRelation ? "semantic" : "canvas" },
edges,
);
if (source.kind === "model_provider" && target.kind === "model_config") {
const next = { ...target, provider: source.name };
return { edges: nextEdges, configs: { ...configs, [target.id]: next } };
}
if (source.kind === "model_config" && target.kind === "llm") {
const next = { ...target, model_alias: source.name };
return { edges: nextEdges, configs: { ...configs, [target.id]: next } };
}
if (
source.kind === "sampler" &&
source.sampler_type === "datetime" &&
target.kind === "sampler" &&
target.sampler_type === "timedelta"
) {
const next = {
...target,
// biome-ignore lint/style/useNamingConvention: api schema
reference_column_name: source.name,
};
return { edges: nextEdges, configs: { ...configs, [target.id]: next } };
}
if (isLlmConfig(target) && source.kind !== "model_provider" && source.kind !== "model_config") {
const ref = `{{ ${source.name} }}`;
const next = {
...target,
prompt: buildTemplateWithRef(target.prompt ?? "", ref),
};
return { edges: nextEdges, configs: { ...configs, [target.id]: next } };
}
if (isExpressionConfig(target) && source.kind !== "model_provider" && source.kind !== "model_config") {
const ref = `{{ ${source.name} }}`;
const next = {
...target,
expr: buildTemplateWithRef(target.expr ?? "", ref),
};
return { edges: nextEdges, configs: { ...configs, [target.id]: next } };
}
if (isSubcategoryConfig(target) && isCategoryConfig(source)) {
const next = syncSubcategoryMapping(target, source);
return { edges: nextEdges, configs: { ...configs, [target.id]: next } };
}
return { edges: nextEdges };
}