diff --git a/studio/frontend/src/features/canvas-lab/blocks/registry.tsx b/studio/frontend/src/features/canvas-lab/blocks/registry.tsx index ac1749dab7..a50a779ce0 100644 --- a/studio/frontend/src/features/canvas-lab/blocks/registry.tsx +++ b/studio/frontend/src/features/canvas-lab/blocks/registry.tsx @@ -9,6 +9,8 @@ import { FunctionIcon, Parabola02Icon, PencilEdit02Icon, + Plant01Icon, + Shield02Icon, Tag01Icon, TagsIcon, UserAccountIcon, @@ -18,10 +20,14 @@ import type { LlmType, NodeConfig, SamplerConfig, SamplerType } from "../types"; import { makeExpressionConfig, makeLlmConfig, + makeModelConfig, + makeModelProviderConfig, makeSamplerConfig, } from "../utils"; import { ExpressionDialog } from "../dialogs/expression/expression-dialog"; import { LlmDialog } from "../dialogs/llm/llm-dialog"; +import { ModelConfigDialog } from "../dialogs/models/model-config-dialog"; +import { ModelProviderDialog } from "../dialogs/models/model-provider-dialog"; import { CategoryDialog } from "../dialogs/samplers/category-dialog"; import { DatetimeDialog } from "../dialogs/samplers/datetime-dialog"; import { GaussianDialog } from "../dialogs/samplers/gaussian-dialog"; @@ -31,7 +37,12 @@ import { UniformDialog } from "../dialogs/samplers/uniform-dialog"; import { UuidDialog } from "../dialogs/samplers/uuid-dialog"; export type BlockKind = "sampler" | "llm" | "expression"; -export type BlockType = SamplerType | LlmType | "expression"; +export type BlockType = + | SamplerType + | LlmType + | "expression" + | "model_provider" + | "model_config"; type IconType = typeof CodeIcon; @@ -249,6 +260,36 @@ const BLOCK_DEFINITIONS: BlockDefinition[] = [ /> ) : null, }, + { + kind: "llm", + type: "model_provider", + title: "Model Provider", + description: "Configure API endpoint + key.", + icon: Shield02Icon, + createConfig: (id, existing) => makeModelProviderConfig(id, existing), + renderDialog: ({ config, onUpdate }) => + config.kind === "model_provider" ? ( + onUpdate(config.id, patch)} + /> + ) : null, + }, + { + kind: "llm", + type: "model_config", + title: "Model Config", + description: "Alias + model + inference params.", + icon: Plant01Icon, + createConfig: (id, existing) => makeModelConfig(id, existing), + renderDialog: ({ config, onUpdate }) => + config.kind === "model_config" ? ( + onUpdate(config.id, patch)} + /> + ) : null, + }, { kind: "expression", type: "expression", @@ -297,6 +338,12 @@ export function getBlockDefinitionForConfig( if (config.kind === "llm") { return getBlockDefinition("llm", config.llm_type); } + if (config.kind === "model_provider") { + return getBlockDefinition("llm", "model_provider"); + } + if (config.kind === "model_config") { + return getBlockDefinition("llm", "model_config"); + } return getBlockDefinition("expression", "expression"); } diff --git a/studio/frontend/src/features/canvas-lab/canvas-lab-page.tsx b/studio/frontend/src/features/canvas-lab/canvas-lab-page.tsx index a726e027d7..c841b9a816 100644 --- a/studio/frontend/src/features/canvas-lab/canvas-lab-page.tsx +++ b/studio/frontend/src/features/canvas-lab/canvas-lab-page.tsx @@ -29,7 +29,7 @@ import { importCanvasPayload } from "./utils/import"; import { buildCanvasPayload } from "./utils/payload"; const NODE_TYPES: NodeTypes = { builder: CanvasNode }; -const EDGE_TYPES: EdgeTypes = { canvas: CanvasEdge }; +const EDGE_TYPES: EdgeTypes = { canvas: CanvasEdge, semantic: CanvasEdge }; type LayoutControlsProps = { direction: "LR" | "TB"; @@ -77,6 +77,8 @@ export function CanvasLabPage(): ReactElement { onConnect, addSamplerNode, addLlmNode, + addModelProviderNode, + addModelConfigNode, addExpressionNode, openConfig, updateConfig, @@ -100,6 +102,8 @@ export function CanvasLabPage(): ReactElement { onConnect: state.onConnect, addSamplerNode: state.addSamplerNode, addLlmNode: state.addLlmNode, + addModelProviderNode: state.addModelProviderNode, + addModelConfigNode: state.addModelConfigNode, addExpressionNode: state.addExpressionNode, openConfig: state.openConfig, updateConfig: state.updateConfig, @@ -306,6 +310,8 @@ export function CanvasLabPage(): ReactElement { onViewChange={setSheetView} onAddSampler={addSamplerNode} onAddLlm={addLlmNode} + onAddModelProvider={addModelProviderNode} + onAddModelConfig={addModelConfigNode} onAddExpression={addExpressionNode} /> diff --git a/studio/frontend/src/features/canvas-lab/components/block-sheet.tsx b/studio/frontend/src/features/canvas-lab/components/block-sheet.tsx index 4e19b6e28e..fb3b2cceba 100644 --- a/studio/frontend/src/features/canvas-lab/components/block-sheet.tsx +++ b/studio/frontend/src/features/canvas-lab/components/block-sheet.tsx @@ -25,6 +25,8 @@ type BlockSheetProps = { onViewChange: (view: SheetView) => void; onAddSampler: (type: SamplerType) => void; onAddLlm: (type: LlmType) => void; + onAddModelProvider: () => void; + onAddModelConfig: () => void; onAddExpression: () => void; }; @@ -92,6 +94,8 @@ export function BlockSheet({ onViewChange, onAddSampler, onAddLlm, + onAddModelProvider, + onAddModelConfig, onAddExpression, }: BlockSheetProps): ReactElement { const title = getSheetTitle(view); @@ -157,7 +161,13 @@ export function BlockSheet({ if (item.kind === "sampler") { onAddSampler(item.type as SamplerType); } else if (item.kind === "llm") { - onAddLlm(item.type as LlmType); + if (item.type === "model_provider") { + onAddModelProvider(); + } else if (item.type === "model_config") { + onAddModelConfig(); + } else { + onAddLlm(item.type as LlmType); + } } else { onAddExpression(); } diff --git a/studio/frontend/src/features/canvas-lab/components/canvas-edge.tsx b/studio/frontend/src/features/canvas-lab/components/canvas-edge.tsx index 6ae2408054..d56906b707 100644 --- a/studio/frontend/src/features/canvas-lab/components/canvas-edge.tsx +++ b/studio/frontend/src/features/canvas-lab/components/canvas-edge.tsx @@ -10,6 +10,7 @@ export const CanvasEdge = memo(function CanvasEdge({ sourcePosition, targetPosition, style, + type, }: EdgeProps): JSX.Element { const [path] = getSmoothStepPath({ sourceX, @@ -22,5 +23,10 @@ export const CanvasEdge = memo(function CanvasEdge({ offset: 16, }); - return ; + const nextStyle = + type === "semantic" + ? { ...style, strokeDasharray: "4 4" } + : style; + + return ; }); diff --git a/studio/frontend/src/features/canvas-lab/components/canvas-node.tsx b/studio/frontend/src/features/canvas-lab/components/canvas-node.tsx index 60283a6bea..0a4b19824c 100644 --- a/studio/frontend/src/features/canvas-lab/components/canvas-node.tsx +++ b/studio/frontend/src/features/canvas-lab/components/canvas-node.tsx @@ -10,6 +10,8 @@ import { FunctionIcon, Parabola02Icon, PencilEdit02Icon, + Plant01Icon, + Shield02Icon, Tag01Icon, TagsIcon, UserAccountIcon, @@ -36,6 +38,12 @@ const NODE_META = { expression: { tone: "bg-sky-50 text-sky-600 border-sky-100", }, + model_provider: { + tone: "bg-amber-50 text-amber-600 border-amber-100", + }, + model_config: { + tone: "bg-indigo-50 text-indigo-600 border-indigo-100", + }, } as const; const SAMPLER_ICONS: Record = { @@ -69,6 +77,12 @@ function resolveNodeIcon( if (kind === "expression") { return FunctionIcon; } + if (kind === "model_provider") { + return Shield02Icon; + } + if (kind === "model_config") { + return Plant01Icon; + } return DiceFaces03Icon; } diff --git a/studio/frontend/src/features/canvas-lab/dialogs/models/model-config-dialog.tsx b/studio/frontend/src/features/canvas-lab/dialogs/models/model-config-dialog.tsx new file mode 100644 index 0000000000..0848ac0ee1 --- /dev/null +++ b/studio/frontend/src/features/canvas-lab/dialogs/models/model-config-dialog.tsx @@ -0,0 +1,121 @@ +import { Checkbox } from "@/components/ui/checkbox"; +import { Input } from "@/components/ui/input"; +import type { ReactElement } from "react"; +import type { ModelConfig } from "../../types"; +import { useCanvasLabStore } from "../../stores/canvas-lab"; +import { NameField } from "../shared/name-field"; + +type ModelConfigDialogProps = { + config: ModelConfig; + onUpdate: (patch: Partial) => void; +}; + +export function ModelConfigDialog({ + config, + onUpdate, +}: ModelConfigDialogProps): ReactElement { + const providerOptions = useCanvasLabStore((state) => + Object.values(state.configs) + .filter((item) => item.kind === "model_provider") + .map((item) => item.name), + ); + const modelId = `${config.id}-model`; + const providerId = `${config.id}-provider`; + const providerListId = `${config.id}-provider-list`; + const tempId = `${config.id}-temperature`; + const topPId = `${config.id}-top-p`; + const maxTokensId = `${config.id}-max-tokens`; + const updateField = ( + key: K, + value: ModelConfig[K], + ) => { + onUpdate({ [key]: value } as Partial); + }; + + return ( +
+ onUpdate({ name: value })} + /> +
+ + updateField("model", event.target.value)} + /> +
+
+ + updateField("provider", event.target.value)} + /> + + {providerOptions.map((provider) => ( + +
+
+ +
+ + updateField("inference_temperature", event.target.value) + } + /> + + updateField("inference_top_p", event.target.value) + } + /> + + updateField("inference_max_tokens", event.target.value) + } + /> +
+
+ +
+ ); +} diff --git a/studio/frontend/src/features/canvas-lab/dialogs/models/model-provider-dialog.tsx b/studio/frontend/src/features/canvas-lab/dialogs/models/model-provider-dialog.tsx new file mode 100644 index 0000000000..182184fd25 --- /dev/null +++ b/studio/frontend/src/features/canvas-lab/dialogs/models/model-provider-dialog.tsx @@ -0,0 +1,130 @@ +import { Input } from "@/components/ui/input"; +import { Textarea } from "@/components/ui/textarea"; +import type { ReactElement } from "react"; +import type { ModelProviderConfig } from "../../types"; +import { NameField } from "../shared/name-field"; + +type ModelProviderDialogProps = { + config: ModelProviderConfig; + onUpdate: (patch: Partial) => void; +}; + +export function ModelProviderDialog({ + config, + onUpdate, +}: ModelProviderDialogProps): ReactElement { + const endpointId = `${config.id}-endpoint`; + const providerTypeId = `${config.id}-provider-type`; + const apiKeyEnvId = `${config.id}-api-key-env`; + const apiKeyId = `${config.id}-api-key`; + const extraHeadersId = `${config.id}-extra-headers`; + const extraBodyId = `${config.id}-extra-body`; + const updateField = ( + key: K, + value: ModelProviderConfig[K], + ) => { + onUpdate({ [key]: value } as Partial); + }; + + return ( +
+ onUpdate({ name: value })} + /> +
+ + + updateField("provider_type", event.target.value) + } + /> +
+
+ + updateField("endpoint", event.target.value)} + /> +
+
+ + updateField("api_key_env", event.target.value)} + /> +
+
+ + updateField("api_key", event.target.value)} + /> +
+
+ +