Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 18 additions & 7 deletions apps/app/src/components/chat/message-list.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,7 @@ import {
getActiveToolLabel,
} from "@/lib/tool-activity"
import { cn } from "@/lib/utils"
import type { ToolInvocationKnownServer } from "@/lib/tool-invocation-origin"
import { groupMessages, isMessageGroup, getLastTextPart, getAssistantRenderGroups, getFileTitle, getMediaBadge, getMessageCreated, formatMessageTimestamp, type UIMessageWithIndex, getMessagesText, getSafeFileDownloadUrl } from "./utils"

const SEARCH_HIGHLIGHT_MARK_CLASS = "rounded px-0.5 bg-amber-4/70 text-current"
Expand All @@ -105,6 +106,7 @@ function MessageTimestamp({ message, className }: { message: UIMessage; classNam

interface ToolMessageProps {
part: ToolUIPart | DynamicToolUIPart
knownServers?: readonly ToolInvocationKnownServer[]
}

/**
Expand All @@ -131,11 +133,11 @@ class ToolMessage extends React.Component<ToolMessageProps, { failed: boolean }>
<div className="text-xs text-muted-foreground">Tool step unavailable</div>
)
}
return <ToolMessageInner part={this.props.part} />
return <ToolMessageInner part={this.props.part} knownServers={this.props.knownServers} />
}
}

const ToolMessageInner = ({ part }: ToolMessageProps) => {
const ToolMessageInner = ({ part, knownServers }: ToolMessageProps) => {
if (isBashToolPart(part)) {
return <BashTool part={part} />
}
Expand Down Expand Up @@ -192,7 +194,7 @@ const ToolMessageInner = ({ part }: ToolMessageProps) => {
return <EnvVarRequestTool part={part} />
}

return <Tool toolPart={part} />
return <Tool toolPart={part} knownServers={knownServers} />
}

const isEmptyMessage = (message: UIMessage): boolean => message.parts.length === 0
Expand Down Expand Up @@ -334,10 +336,11 @@ type AssistantMessageProps = {
isLastMessage: boolean
isStreaming: boolean
isLastStep: boolean
knownServers?: readonly ToolInvocationKnownServer[]
}

const AssistantMessage = React.memo(
({ message }: AssistantMessageProps) => {
({ message, knownServers }: AssistantMessageProps) => {
const { showThinking, highlightQuery } = useMessageList()
const assistantRenderGroups = React.useMemo(
() => getAssistantRenderGroups(message.parts, showThinking),
Expand Down Expand Up @@ -387,7 +390,7 @@ const AssistantMessage = React.memo(

return (
<div key={`tool-${index}`} className="w-full">
<ToolMessage part={group.part} />
<ToolMessage part={group.part} knownServers={knownServers} />
</div>
)
})}
Expand Down Expand Up @@ -567,10 +570,11 @@ type MessageComponentProps = {
isLastMessage: boolean
isStreaming: boolean
isLastStep: boolean
knownServers?: readonly ToolInvocationKnownServer[]
}

const MessageComponent = React.memo(
({ message, isLastMessage, isStreaming, isLastStep }: MessageComponentProps) => {
({ message, isLastMessage, isStreaming, isLastStep, knownServers }: MessageComponentProps) => {
if (isSessionErrorMessage(message)) {
return <ErrorMessage error={getMessagesText([message]) || "Session failed"} />
}
Expand All @@ -591,6 +595,7 @@ const MessageComponent = React.memo(
isLastMessage={isLastMessage}
isStreaming={isStreaming}
isLastStep={isLastStep}
knownServers={knownServers}
/>
)
}
Expand Down Expand Up @@ -731,12 +736,14 @@ interface AssistantMessageGroupProps {
items: UIMessageWithIndex[]
messages: UIMessage[]
isStreaming: boolean
knownServers?: readonly ToolInvocationKnownServer[]
}

function MessageGroup({
items,
messages,
isStreaming,
knownServers,
}: AssistantMessageGroupProps) {
const { onRevertToUserMessage, onForkAtMessage } = useMessageList()
const lastItem = items[items.length - 1]
Expand Down Expand Up @@ -786,6 +793,7 @@ function MessageGroup({
isLastMessage={isLastMessage}
isStreaming={isLastMessage && isStreaming}
isLastStep={groupIndex === items.length - 1}
knownServers={knownServers}
/>
<MessageArtifacts message={item.message} />
</div>
Expand Down Expand Up @@ -842,9 +850,10 @@ interface MessageListProps {
messages: UIMessage[]
status: ThreadStatus
retryStatus?: RetryStatus | null
knownMcpServers?: readonly ToolInvocationKnownServer[]
}

export function MessageList({ messages, status, retryStatus }: MessageListProps) {
export function MessageList({ messages, status, retryStatus, knownMcpServers }: MessageListProps) {
const isStreaming = status === "streaming" || status === "retrying"
const items = React.useMemo(() => groupMessages(messages, status), [messages, status]);
const error = useSessionErrorMessage();
Expand All @@ -865,6 +874,7 @@ export function MessageList({ messages, status, retryStatus }: MessageListProps)
items={item.messages}
messages={messages}
isStreaming={isStreaming}
knownServers={knownMcpServers}
/>
)
}
Expand All @@ -880,6 +890,7 @@ export function MessageList({ messages, status, retryStatus }: MessageListProps)
isLastMessage={isLastMessage}
isStreaming={isLastMessage && isStreaming}
isLastStep={isLastStep}
knownServers={knownMcpServers}
/>
<MessageArtifacts message={item.message} />
</div>
Expand Down
14 changes: 12 additions & 2 deletions apps/app/src/components/ui/tool.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ import {
Wrench,
} from "lucide-react"
import type { DynamicToolUIPart, ToolUIPart } from "ai"
import { toolInvocationOrigin, type ToolInvocationKnownServer } from "@/lib/tool-invocation-origin"

function toolIcon(part: ToolPart) {
const name = part.type === "dynamic-tool" ? part.toolName : part.type
Expand Down Expand Up @@ -56,6 +57,7 @@ export type ToolPart = ToolUIPart | DynamicToolUIPart
export type ToolProps = {
title?: string
toolPart: ToolPart
knownServers?: readonly ToolInvocationKnownServer[]
defaultOpen?: boolean
className?: string
}
Expand Down Expand Up @@ -116,11 +118,14 @@ function DiffLines({ diff }: { diff: string }) {
)
}

const Tool = ({ title, toolPart, defaultOpen = false, className }: ToolProps) => {
const Tool = ({ title, toolPart, knownServers, defaultOpen = false, className }: ToolProps) => {
const { state, input } = toolPart
const inFlight = isToolPartInFlight(toolPart)
const isError = state === "output-error"
const label = title ?? getToolActivityLabel(toolPart)
const origin = toolPart.type === "dynamic-tool"
? toolInvocationOrigin(toolPart.toolName, input, knownServers)
: null
const label = title ?? (origin?.connectionName ? origin.displayTool : getToolActivityLabel(toolPart))
const hasInput = input !== null && input !== undefined
const hasOutput = "output" in toolPart && toolPart.output !== undefined
const inputDiff = getInputDiff(input)
Expand All @@ -143,6 +148,11 @@ const Tool = ({ title, toolPart, defaultOpen = false, className }: ToolProps) =>
</span>
<ChevronDown className="absolute size-4 opacity-0 transition-opacity group-hover:opacity-100 group-data-panel-open:rotate-180" />
</span>
{origin?.connectionName ? (
<span className="shrink-0 rounded-md border border-border bg-muted px-1.5 py-0.5 text-[10px] font-medium text-muted-foreground">
{origin.connectionName}
</span>
) : null}
<span className="min-w-0 truncate">{label}</span>
{isError ? (
<span className="text-destructive shrink-0 text-xs">failed</span>
Expand Down
193 changes: 193 additions & 0 deletions apps/app/src/lib/tool-invocation-origin.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,193 @@
export type ToolInvocationKnownServer = {
id?: string | null;
name: string;
displayName?: string | null;
};

export type ToolInvocationOrigin = {
connectionName?: string;
displayTool: string;
};

type DirectToolOrigin = ToolInvocationOrigin & {
serverName: string;
};

const OPENWORK_CLOUD_SERVER_NAME = "openwork-cloud";
const OPENWORK_CLOUD_LABEL = "OpenWork Cloud";
const EXECUTE_CAPABILITY_TOOL_NAME = "execute_capability";

const builtInToolNames = new Set([
"apply_patch",
"bash",
"edit",
"env_var_request",
"glob",
"grep",
"lsp",
"question",
"read",
"request_env_var",
"skill",
"task",
"todowrite",
"webfetch",
"websearch",
"write",
]);

const nonConnectionPrefixes = new Set([
"create",
"delete",
"diagnostic",
"diagnostics",
"execute",
"fetch",
"get",
"list",
"lookup",
"mutate",
"read",
"search",
"summarize",
"synthetic",
"update",
"write",
]);

const wordLabels: Record<string, string> = {
api: "API",
mcp: "MCP",
oauth: "OAuth",
openwork: "OpenWork",
servicenow: "ServiceNow",
ui: "UI",
};

function formatNameWord(word: string) {
const lower = word.toLowerCase();
const special = wordLabels[lower];
if (special) return special;
return lower.charAt(0).toUpperCase() + lower.slice(1);
}

function formatConnectionName(value: string) {
const words = value.replace(/[_-]+/g, " ").trim().split(/\s+/).filter(Boolean);
return words.length > 0 ? words.map(formatNameWord).join(" ") : value;
}

function serverCandidates(server: ToolInvocationKnownServer) {
const candidates: string[] = [];
const name = server.name.trim();
const id = server.id?.trim() ?? "";
if (name) candidates.push(name);
if (id && id !== name) candidates.push(id);
return candidates;
}

function candidateLength(server: ToolInvocationKnownServer) {
return Math.max(...serverCandidates(server).map((candidate) => candidate.length), 0);
}

function displayNameForServer(server: ToolInvocationKnownServer) {
const explicit = server.displayName?.trim();
if (explicit) return explicit;
return formatConnectionName(server.name);
}

function sortedKnownServers(knownServers: readonly ToolInvocationKnownServer[]) {
return [...knownServers].sort((left, right) => candidateLength(right) - candidateLength(left));
}

function resolveKnownConnectionName(id: string, knownServers: readonly ToolInvocationKnownServer[]) {
const trimmed = id.trim();
if (!trimmed) return null;
const match = knownServers.find((server) => {
const serverId = server.id?.trim() ?? "";
return serverId === trimmed || server.name.trim() === trimmed;
});
return match ? displayNameForServer(match) : null;
}

function directToolOrigin(toolName: string, knownServers: readonly ToolInvocationKnownServer[]): DirectToolOrigin | null {
for (const server of sortedKnownServers(knownServers)) {
for (const candidate of serverCandidates(server)) {
const prefix = `${candidate}_`;
if (!toolName.startsWith(prefix)) continue;
const displayTool = toolName.slice(prefix.length);
if (!displayTool) continue;
return {
connectionName: displayNameForServer(server),
displayTool,
serverName: server.name.trim(),
};
}
}

const openworkPrefix = `${OPENWORK_CLOUD_SERVER_NAME}_`;
if (toolName.startsWith(openworkPrefix)) {
const displayTool = toolName.slice(openworkPrefix.length);
if (displayTool) {
return {
connectionName: OPENWORK_CLOUD_LABEL,
displayTool,
serverName: OPENWORK_CLOUD_SERVER_NAME,
};
}
}

return null;
}

function fallbackDirectToolOrigin(toolName: string): DirectToolOrigin | null {
const separator = toolName.indexOf("_");
if (separator <= 0 || separator === toolName.length - 1) return null;
const serverName = toolName.slice(0, separator).trim();
const displayTool = toolName.slice(separator + 1).trim();
if (!serverName || !displayTool || nonConnectionPrefixes.has(serverName.toLowerCase())) return null;
return {
connectionName: formatConnectionName(serverName),
displayTool,
serverName,
};
}

function namedArgument(value: unknown) {
if (typeof value !== "object" || value === null || Array.isArray(value) || !("name" in value)) return null;
return typeof value.name === "string" ? value.name : null;
}

function mcpCapabilityNameParts(value: string) {
const match = value.match(/^mcp:([^:]+):(.+)$/);
const connectionId = match?.[1]?.trim() ?? "";
const toolName = match?.[2]?.trim() ?? "";
if (!connectionId || !toolName) return null;
return { connectionId, toolName };
}

export function toolInvocationOrigin(
toolName: string,
args?: unknown,
knownServers: readonly ToolInvocationKnownServer[] = [],
): ToolInvocationOrigin {
if (builtInToolNames.has(toolName)) return { displayTool: toolName };

const direct = directToolOrigin(toolName, knownServers) ?? fallbackDirectToolOrigin(toolName);
if (!direct) return { displayTool: toolName };

if (direct.serverName === OPENWORK_CLOUD_SERVER_NAME && direct.displayTool === EXECUTE_CAPABILITY_TOOL_NAME) {
const parts = mcpCapabilityNameParts(namedArgument(args) ?? "");
if (parts) {
const connectionName = resolveKnownConnectionName(parts.connectionId, knownServers) ?? parts.connectionId;
return {
connectionName: `${OPENWORK_CLOUD_LABEL} → ${connectionName}`,
displayTool: parts.toolName,
};
}
}

return {
connectionName: direct.connectionName,
displayTool: direct.displayTool,
};
}
Loading
Loading