Merge pull request #3917 from code-yeongyu/fix/anthropic-prefill-recovery
fix(plugin): guard assistant prefill message tails
This commit is contained in:
@@ -7,10 +7,13 @@ import type { CreatedHooks } from "../create-hooks"
|
|||||||
type TestPart = {
|
type TestPart = {
|
||||||
type: string
|
type: string
|
||||||
id?: string
|
id?: string
|
||||||
|
sessionID?: string
|
||||||
|
messageID?: string
|
||||||
callID?: string
|
callID?: string
|
||||||
tool_use_id?: string
|
tool_use_id?: string
|
||||||
content?: string
|
content?: string
|
||||||
text?: string
|
text?: string
|
||||||
|
synthetic?: boolean
|
||||||
}
|
}
|
||||||
|
|
||||||
type TestMessage = {
|
type TestMessage = {
|
||||||
@@ -38,7 +41,7 @@ function makeHooks(overrides: {
|
|||||||
contextInjectorMessagesTransform: overrides.contextInjector ? makeHook(overrides.contextInjector) : undefined,
|
contextInjectorMessagesTransform: overrides.contextInjector ? makeHook(overrides.contextInjector) : undefined,
|
||||||
thinkingBlockValidator: overrides.thinkingBlock ? makeHook(overrides.thinkingBlock) : undefined,
|
thinkingBlockValidator: overrides.thinkingBlock ? makeHook(overrides.thinkingBlock) : undefined,
|
||||||
toolPairValidator: overrides.toolPair ? makeHook(overrides.toolPair) : undefined,
|
toolPairValidator: overrides.toolPair ? makeHook(overrides.toolPair) : undefined,
|
||||||
} as unknown as CreatedHooks
|
} as CreatedHooks
|
||||||
}
|
}
|
||||||
|
|
||||||
async function runHandler(
|
async function runHandler(
|
||||||
@@ -157,6 +160,25 @@ describe("createMessagesTransformHandler", () => {
|
|||||||
//#when / #then
|
//#when / #then
|
||||||
await runHandler(hooks, [])
|
await runHandler(hooks, [])
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it("appends a synthetic user turn when transformed messages end with assistant prefill", async () => {
|
||||||
|
//#given
|
||||||
|
const messages: TestMessage[] = [
|
||||||
|
{ info: { role: "user" }, parts: [{ type: "text", text: "work on this" }] },
|
||||||
|
{ info: { role: "assistant" }, parts: [{ type: "text", text: "partial assistant tail" }] },
|
||||||
|
]
|
||||||
|
|
||||||
|
//#when
|
||||||
|
await runHandler(makeHooks({}), messages)
|
||||||
|
|
||||||
|
//#then
|
||||||
|
expect(messages.at(-1)?.info).toMatchObject({ role: "user" })
|
||||||
|
expect(messages.at(-1)?.parts[0]).toMatchObject({
|
||||||
|
type: "text",
|
||||||
|
text: "[internal] Continue from the previous assistant state.",
|
||||||
|
synthetic: true,
|
||||||
|
})
|
||||||
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
function createRealToolPairValidator(): TransformHook {
|
function createRealToolPairValidator(): TransformHook {
|
||||||
|
|||||||
@@ -3,12 +3,75 @@ import type { Message, Part } from "@opencode-ai/sdk"
|
|||||||
import { log } from "../shared/logger"
|
import { log } from "../shared/logger"
|
||||||
import type { CreatedHooks } from "../create-hooks"
|
import type { CreatedHooks } from "../create-hooks"
|
||||||
|
|
||||||
|
const ASSISTANT_PREFILL_RECOVERY_TEXT = "[internal] Continue from the previous assistant state."
|
||||||
|
|
||||||
type MessageWithParts = {
|
type MessageWithParts = {
|
||||||
info: Message
|
info: Message
|
||||||
parts: Part[]
|
parts: Part[]
|
||||||
}
|
}
|
||||||
|
|
||||||
type MessagesTransformOutput = { messages: MessageWithParts[] }
|
type MessagesTransformOutput = { messages: MessageWithParts[] }
|
||||||
|
type UserMessageInfo = Extract<Message, { role: "user" }>
|
||||||
|
|
||||||
|
function getSessionID(message: MessageWithParts): string | undefined {
|
||||||
|
return message.info.sessionID
|
||||||
|
}
|
||||||
|
|
||||||
|
function findLastUserMessage(messages: MessageWithParts[]): UserMessageInfo | undefined {
|
||||||
|
for (let index = messages.length - 1; index >= 0; index -= 1) {
|
||||||
|
const message = messages[index]
|
||||||
|
if (message?.info.role === "user") {
|
||||||
|
return message.info
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return undefined
|
||||||
|
}
|
||||||
|
|
||||||
|
function createAssistantPrefillRecoveryMessage(
|
||||||
|
lastAssistantMessage: MessageWithParts,
|
||||||
|
messages: MessageWithParts[],
|
||||||
|
): MessageWithParts {
|
||||||
|
const lastUserMessage = findLastUserMessage(messages)
|
||||||
|
const sessionID = getSessionID(lastAssistantMessage) ?? lastUserMessage?.sessionID ?? ""
|
||||||
|
const messageID = `${lastAssistantMessage.info.id}_prefill_recovery`
|
||||||
|
const model = lastUserMessage?.model ?? {
|
||||||
|
providerID: "internal",
|
||||||
|
modelID: "assistant-prefill-guard",
|
||||||
|
}
|
||||||
|
|
||||||
|
return {
|
||||||
|
info: {
|
||||||
|
id: messageID,
|
||||||
|
sessionID,
|
||||||
|
role: "user",
|
||||||
|
time: { created: Date.now() },
|
||||||
|
agent: lastUserMessage?.agent ?? "internal",
|
||||||
|
model,
|
||||||
|
...(lastUserMessage?.system ? { system: lastUserMessage.system } : {}),
|
||||||
|
...(lastUserMessage?.tools ? { tools: lastUserMessage.tools } : {}),
|
||||||
|
},
|
||||||
|
parts: [
|
||||||
|
{
|
||||||
|
id: `${messageID}_text`,
|
||||||
|
sessionID,
|
||||||
|
messageID,
|
||||||
|
type: "text",
|
||||||
|
text: ASSISTANT_PREFILL_RECOVERY_TEXT,
|
||||||
|
synthetic: true,
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function ensureUserTurnAfterAssistantTail(output: MessagesTransformOutput): void {
|
||||||
|
const lastMessage = output.messages.at(-1)
|
||||||
|
if (!lastMessage || lastMessage.info.role !== "assistant") {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
output.messages.push(createAssistantPrefillRecoveryMessage(lastMessage, output.messages))
|
||||||
|
}
|
||||||
|
|
||||||
async function runMessagesTransformHookSafely<I, O>(
|
async function runMessagesTransformHookSafely<I, O>(
|
||||||
hookName: string,
|
hookName: string,
|
||||||
@@ -79,5 +142,7 @@ export function createMessagesTransformHandler(args: {
|
|||||||
input,
|
input,
|
||||||
output,
|
output,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
ensureUserTurnAfterAssistantTail(output)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user