@shanepadgett/tau-agent 0.6.0 → 0.7.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/extensions/auto-name/index.ts +1 -1
- package/extensions/commit/commit-plan.ts +12 -5
- package/extensions/context/index.ts +23 -20
- package/extensions/context/sync.ts +28 -18
- package/extensions/image-gen/README.md +4 -4
- package/extensions/image-gen/client.ts +126 -107
- package/extensions/image-gen/index.ts +44 -35
- package/extensions/tau-help/help.md +4 -0
- package/extensions/xai/README.md +7 -0
- package/extensions/xai/auth.ts +40 -0
- package/extensions/xai/constants.ts +11 -0
- package/extensions/xai/index.ts +38 -0
- package/extensions/xai/oauth.ts +342 -0
- package/extensions/xai/payload.ts +68 -0
- package/package.json +2 -2
- package/shared/model-fallback/index.ts +17 -12
|
@@ -87,7 +87,7 @@ async function runAutoName(
|
|
|
87
87
|
): Promise<void> {
|
|
88
88
|
const ui = ctx.ui;
|
|
89
89
|
try {
|
|
90
|
-
const candidates = await resolveCandidates(ctx, AUTO_NAME_MODELS);
|
|
90
|
+
const candidates = await resolveCandidates(ctx, AUTO_NAME_MODELS, true);
|
|
91
91
|
const result = await generateToolValidated(
|
|
92
92
|
{ ui, signal: controller.signal },
|
|
93
93
|
candidates,
|
|
@@ -55,7 +55,7 @@ export async function generatePlan(
|
|
|
55
55
|
const prompt = buildPlanPrompt(evidence, previousPlan, regenerationNote);
|
|
56
56
|
return generateToolValidated(
|
|
57
57
|
ctx,
|
|
58
|
-
await resolveCandidates(ctx, COMMIT_MODELS),
|
|
58
|
+
await resolveCandidates(ctx, COMMIT_MODELS, true),
|
|
59
59
|
prompt,
|
|
60
60
|
COMMIT_PLAN_TOOL,
|
|
61
61
|
(input) => commitGroupsFromToolInput(input, evidence.files),
|
|
@@ -83,10 +83,17 @@ export async function regenerateMessage(
|
|
|
83
83
|
): Promise<string> {
|
|
84
84
|
const selected = evidence.files.filter((file) => files.includes(file.path));
|
|
85
85
|
const prompt = buildMessagePrompt(evidence, selected, previousPlan, selectedGroupId, regenerationNote);
|
|
86
|
-
return generateValidated(
|
|
87
|
-
|
|
88
|
-
|
|
89
|
-
|
|
86
|
+
return generateValidated(
|
|
87
|
+
ctx,
|
|
88
|
+
await resolveCandidates(ctx, COMMIT_MODELS, true),
|
|
89
|
+
prompt,
|
|
90
|
+
requireCommitMessage,
|
|
91
|
+
undefined,
|
|
92
|
+
{
|
|
93
|
+
statusKey: "commit",
|
|
94
|
+
notifyOnFallback: true,
|
|
95
|
+
},
|
|
96
|
+
);
|
|
90
97
|
}
|
|
91
98
|
|
|
92
99
|
export function requireCommitMessage(rawMessage: string): string {
|
|
@@ -89,39 +89,42 @@ export default function contextExtension(pi: ExtensionAPI): void {
|
|
|
89
89
|
});
|
|
90
90
|
|
|
91
91
|
pi.registerTool(
|
|
92
|
-
defineTool<typeof contextSyncParams, ContextSyncDetails>({
|
|
92
|
+
defineTool<typeof contextSyncParams, ContextSyncDetails | undefined>({
|
|
93
93
|
name: "context_sync",
|
|
94
94
|
label: "context_sync",
|
|
95
95
|
description: "Synchronize repository context from current Git changes.",
|
|
96
96
|
parameters: contextSyncParams,
|
|
97
|
-
async execute(_id, _params, _signal,
|
|
98
|
-
return compactResult(
|
|
97
|
+
async execute(_id, _params, _signal, onUpdate, ctx) {
|
|
98
|
+
return compactResult(
|
|
99
|
+
await runContextSync(pi, ctx, (status) =>
|
|
100
|
+
onUpdate?.({ content: [{ type: "text", text: status }], details: undefined }),
|
|
101
|
+
),
|
|
102
|
+
);
|
|
99
103
|
},
|
|
100
104
|
renderCall(_args, theme, context) {
|
|
101
105
|
const text = (context.lastComponent as Text | undefined) ?? new Text("", 0, 0);
|
|
102
106
|
text.setText(theme.fg("toolTitle", "context_sync"));
|
|
103
107
|
return text;
|
|
104
108
|
},
|
|
105
|
-
renderResult(result,
|
|
109
|
+
renderResult(result, options, theme, context) {
|
|
106
110
|
const text = (context.lastComponent as Text | undefined) ?? new Text("", 0, 0);
|
|
107
111
|
const details = result.details;
|
|
112
|
+
const output = result.content.map((part) => (part.type === "text" ? part.text : "")).join("");
|
|
108
113
|
text.setText(
|
|
109
|
-
|
|
110
|
-
?
|
|
111
|
-
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
|
|
119
|
-
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
result.content.map((part) => (part.type === "text" ? part.text : "")).join(""),
|
|
124
|
-
)),
|
|
114
|
+
options.isPartial
|
|
115
|
+
? theme.fg("dim", output)
|
|
116
|
+
: context.expanded && details
|
|
117
|
+
? [
|
|
118
|
+
details.summary,
|
|
119
|
+
details.reason,
|
|
120
|
+
...details.changes.map((change) =>
|
|
121
|
+
change.action === "set-entry"
|
|
122
|
+
? `${change.action} ${change.tab}/${change.concept}/${change.entry}: ${change.files.join(", ")}`
|
|
123
|
+
: `${change.action} ${change.tab}/${change.concept}/${change.entry}`,
|
|
124
|
+
),
|
|
125
|
+
...details.changedContextFiles,
|
|
126
|
+
].join("\n")
|
|
127
|
+
: (details?.summary ?? theme.fg("error", output)),
|
|
125
128
|
);
|
|
126
129
|
return text;
|
|
127
130
|
},
|
|
@@ -1,12 +1,13 @@
|
|
|
1
1
|
import { createHash } from "node:crypto";
|
|
2
2
|
import { mkdir, readFile, readdir, rename, rm, stat, writeFile } from "node:fs/promises";
|
|
3
3
|
import { dirname, extname, join, relative, resolve, sep } from "node:path";
|
|
4
|
-
import { Type, type Tool } from "@earendil-works/pi-ai";
|
|
4
|
+
import { Type, type ThinkingLevel, type Tool } from "@earendil-works/pi-ai";
|
|
5
5
|
import { withFileMutationQueue, type ExtensionAPI, type ExtensionContext } from "@earendil-works/pi-coding-agent";
|
|
6
6
|
import { parse, stringify } from "smol-toml";
|
|
7
7
|
import { createGitRunner, loadRepoStatus, type GitRunner } from "../../shared/git.ts";
|
|
8
8
|
import { generateToolValidated, resolveCandidates } from "../../shared/model-fallback/index.ts";
|
|
9
9
|
import { truncAt } from "../../shared/text.ts";
|
|
10
|
+
import { XAI_CHAT_MODEL, XAI_PROVIDER } from "../xai/constants.ts";
|
|
10
11
|
import {
|
|
11
12
|
loadContextEntries,
|
|
12
13
|
normalizeProjectPath,
|
|
@@ -22,19 +23,22 @@ const MAX_TOTAL_EVIDENCE = 64_000;
|
|
|
22
23
|
const MAX_UNTRACKED_BYTES = 12_000;
|
|
23
24
|
const EVIDENCE_CONCURRENCY = 4;
|
|
24
25
|
|
|
26
|
+
const CONTEXT_SYNC_MODELS: ReadonlyArray<{ provider: string; model: string; reasoning: ThinkingLevel }> = [
|
|
27
|
+
{ provider: "openai-codex", model: "gpt-5.6-terra", reasoning: "medium" },
|
|
28
|
+
{ provider: "openai-codex", model: "gpt-5.6-sol", reasoning: "low" },
|
|
29
|
+
{ provider: "anthropic", model: "claude-sonnet-5", reasoning: "low" },
|
|
30
|
+
{ provider: XAI_PROVIDER, model: XAI_CHAT_MODEL, reasoning: "high" },
|
|
31
|
+
];
|
|
32
|
+
|
|
25
33
|
const SUBMIT_TOOL = {
|
|
26
34
|
name: "submit_context_sync",
|
|
27
35
|
description: "Submit the desired context catalog changes.",
|
|
28
|
-
parameters: Type.
|
|
29
|
-
|
|
30
|
-
|
|
31
|
-
{
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
{
|
|
35
|
-
outcome: Type.Literal("apply"),
|
|
36
|
-
reason: Type.String({ minLength: 1 }),
|
|
37
|
-
changes: Type.Array(
|
|
36
|
+
parameters: Type.Object(
|
|
37
|
+
{
|
|
38
|
+
outcome: Type.Union([Type.Literal("no-change"), Type.Literal("apply")]),
|
|
39
|
+
reason: Type.String({ minLength: 1 }),
|
|
40
|
+
changes: Type.Optional(
|
|
41
|
+
Type.Array(
|
|
38
42
|
Type.Union([
|
|
39
43
|
Type.Object(
|
|
40
44
|
{
|
|
@@ -61,10 +65,10 @@ const SUBMIT_TOOL = {
|
|
|
61
65
|
]),
|
|
62
66
|
{ minItems: 1 },
|
|
63
67
|
),
|
|
64
|
-
|
|
65
|
-
|
|
66
|
-
|
|
67
|
-
|
|
68
|
+
),
|
|
69
|
+
},
|
|
70
|
+
{ additionalProperties: false },
|
|
71
|
+
),
|
|
68
72
|
} satisfies Tool;
|
|
69
73
|
|
|
70
74
|
export interface SyncDirtyFile {
|
|
@@ -122,8 +126,13 @@ export interface SyncEvidence {
|
|
|
122
126
|
|
|
123
127
|
let syncQueue = Promise.resolve();
|
|
124
128
|
|
|
125
|
-
export async function runContextSync(
|
|
129
|
+
export async function runContextSync(
|
|
130
|
+
pi: ExtensionAPI,
|
|
131
|
+
ctx: ExtensionContext,
|
|
132
|
+
onStatus?: (status: string) => void | Promise<void>,
|
|
133
|
+
): Promise<ContextSyncDetails> {
|
|
126
134
|
if (!ctx.isProjectTrusted()) throw new Error("Context sync requires a trusted project");
|
|
135
|
+
await onStatus?.("Inspecting repository context");
|
|
127
136
|
const git = createGitRunner(pi, ctx);
|
|
128
137
|
const status = await loadRepoStatus(git);
|
|
129
138
|
if (!status) throw new Error("No Git repository found");
|
|
@@ -132,14 +141,15 @@ export async function runContextSync(pi: ExtensionAPI, ctx: ExtensionContext): P
|
|
|
132
141
|
const prompt = buildContextSyncPrompt(evidence);
|
|
133
142
|
const plan = await generateToolValidated(
|
|
134
143
|
ctx,
|
|
135
|
-
await resolveCandidates(ctx),
|
|
144
|
+
await resolveCandidates(ctx, CONTEXT_SYNC_MODELS, false),
|
|
136
145
|
prompt,
|
|
137
146
|
SUBMIT_TOOL,
|
|
138
147
|
(input) => normalizeContextSyncPlan(input, evidence),
|
|
139
148
|
(error) => `Validation failed: ${error.message}\nCall submit_context_sync once with corrected arguments only.`,
|
|
140
|
-
{ statusKey: "context-sync",
|
|
149
|
+
{ statusKey: "context-sync", onStatus },
|
|
141
150
|
);
|
|
142
151
|
if (plan.outcome === "no-change") return noChange(plan.reason);
|
|
152
|
+
await onStatus?.("Applying context catalog changes");
|
|
143
153
|
return withSyncLock(async () => {
|
|
144
154
|
return applyContextSyncPlan(evidence.root, plan, evidence.entries, async () => {
|
|
145
155
|
const currentEntries = await loadContextEntries(evidence.root);
|
|
@@ -1,11 +1,11 @@
|
|
|
1
1
|
# Image Generation
|
|
2
2
|
|
|
3
|
-
`image_gen` generates raster images and edits up to
|
|
3
|
+
`image_gen` generates raster images and edits up to three local raster images with Grok Imagine. It uses `grok-imagine-image-quality` and saves results under `~/.local/share/tau-agent/images/` by default. Pass an explicit path with the expected image extension when the image should be saved in the current repository or another chosen location.
|
|
4
4
|
|
|
5
|
-
Use `/login` and select
|
|
5
|
+
Use `/login` and select xAI (Grok subscription OAuth) before invoking the tool. No xAI API key is used.
|
|
6
6
|
|
|
7
7
|
Run `/reload` after installing or changing the extension.
|
|
8
8
|
|
|
9
|
-
The model invokes `image_gen` with a prompt.
|
|
9
|
+
The model invokes `image_gen` with a prompt. For edits, it also supplies one to three local PNG, JPEG, or WebP paths. Successful images up to 12 MiB are returned inline for inspection; larger results remain available at the saved path.
|
|
10
10
|
|
|
11
|
-
This extension
|
|
11
|
+
This extension uses xAI's undocumented subscription OAuth access. xAI may change its availability, entitlement rules, or protocol without notice.
|
|
@@ -1,60 +1,34 @@
|
|
|
1
|
-
|
|
2
|
-
|
|
3
|
-
const PNG_SIGNATURE = Buffer.from([0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a]);
|
|
1
|
+
import { XAI_API_BASE_URL, XAI_IMAGE_MODEL } from "../xai/constants.ts";
|
|
2
|
+
|
|
4
3
|
const MAX_ERROR_BODY_BYTES = 8192;
|
|
5
4
|
const MAX_ERROR_MESSAGE_LENGTH = 2000;
|
|
6
|
-
|
|
7
|
-
|
|
8
|
-
|
|
9
|
-
status: number;
|
|
10
|
-
body: ReadableStream<Uint8Array> | null;
|
|
11
|
-
json(): Promise<unknown>;
|
|
12
|
-
}
|
|
13
|
-
|
|
14
|
-
export interface CodexAuth {
|
|
15
|
-
token: string;
|
|
16
|
-
accountId: string;
|
|
17
|
-
}
|
|
5
|
+
const REQUEST_TIMEOUT_MS = 60_000;
|
|
6
|
+
const MAX_ATTEMPTS = 3;
|
|
7
|
+
const RETRY_BASE_DELAY_MS = 500;
|
|
18
8
|
|
|
19
9
|
export interface EditImage {
|
|
20
10
|
mimeType: "image/png" | "image/jpeg" | "image/webp";
|
|
21
11
|
data: string;
|
|
22
12
|
}
|
|
23
13
|
|
|
24
|
-
interface GeneratedImage {
|
|
14
|
+
export interface GeneratedImage {
|
|
25
15
|
bytes: Buffer;
|
|
26
16
|
base64: string;
|
|
27
|
-
mimeType: "
|
|
17
|
+
mimeType: EditImage["mimeType"];
|
|
28
18
|
}
|
|
29
19
|
|
|
30
|
-
|
|
31
|
-
|
|
20
|
+
interface HttpResponse {
|
|
21
|
+
ok: boolean;
|
|
22
|
+
status: number;
|
|
23
|
+
body: ReadableStream<Uint8Array> | null;
|
|
24
|
+
json(): Promise<unknown>;
|
|
32
25
|
}
|
|
33
26
|
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
new Error("The OpenAI Codex credential does not contain a usable ChatGPT account ID. Run /login again.");
|
|
37
|
-
const segments = token.split(".");
|
|
38
|
-
if (segments.length !== 3 || !segments[1]) throw invalidCredential();
|
|
39
|
-
if (!/^[A-Za-z0-9_-]+$/.test(segments[1]) || segments[1].length % 4 === 1) throw invalidCredential();
|
|
40
|
-
|
|
41
|
-
let payload: unknown;
|
|
42
|
-
try {
|
|
43
|
-
const decoded = Buffer.from(segments[1], "base64url");
|
|
44
|
-
if (decoded.toString("base64url") !== segments[1]) throw invalidCredential();
|
|
45
|
-
payload = JSON.parse(decoded.toString("utf8"));
|
|
46
|
-
} catch {
|
|
47
|
-
throw invalidCredential();
|
|
48
|
-
}
|
|
49
|
-
if (!isRecord(payload)) throw invalidCredential();
|
|
50
|
-
const authClaim = payload["https://api.openai.com/auth"];
|
|
51
|
-
if (!isRecord(authClaim)) throw invalidCredential();
|
|
52
|
-
const accountId = authClaim.chatgpt_account_id;
|
|
53
|
-
if (typeof accountId !== "string" || !accountId.trim()) throw invalidCredential();
|
|
54
|
-
return { token, accountId: accountId.trim() };
|
|
27
|
+
function isRecord(value: unknown): value is Record<string, unknown> {
|
|
28
|
+
return typeof value === "object" && value !== null && !Array.isArray(value);
|
|
55
29
|
}
|
|
56
30
|
|
|
57
|
-
async function
|
|
31
|
+
async function boundedError(response: HttpResponse): Promise<string> {
|
|
58
32
|
if (!response.body) return "";
|
|
59
33
|
const reader = response.body.getReader();
|
|
60
34
|
const chunks: Uint8Array[] = [];
|
|
@@ -63,8 +37,7 @@ async function readBoundedError(response: FetchResponse): Promise<string> {
|
|
|
63
37
|
while (length < MAX_ERROR_BODY_BYTES) {
|
|
64
38
|
const result = await reader.read();
|
|
65
39
|
if (result.done) break;
|
|
66
|
-
const
|
|
67
|
-
const chunk = result.value.subarray(0, remaining);
|
|
40
|
+
const chunk = result.value.subarray(0, MAX_ERROR_BODY_BYTES - length);
|
|
68
41
|
chunks.push(chunk);
|
|
69
42
|
length += chunk.length;
|
|
70
43
|
if (chunk.length < result.value.length) break;
|
|
@@ -72,23 +45,21 @@ async function readBoundedError(response: FetchResponse): Promise<string> {
|
|
|
72
45
|
} finally {
|
|
73
46
|
await reader.cancel().catch(() => undefined);
|
|
74
47
|
}
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
return new TextDecoder().decode(bytes);
|
|
48
|
+
return Buffer.concat(
|
|
49
|
+
chunks.map((chunk) => Buffer.from(chunk)),
|
|
50
|
+
length,
|
|
51
|
+
)
|
|
52
|
+
.toString()
|
|
53
|
+
.trim();
|
|
82
54
|
}
|
|
83
55
|
|
|
84
56
|
function serverErrorMessage(body: string, token: string): string {
|
|
85
|
-
let message = body
|
|
57
|
+
let message = body;
|
|
86
58
|
try {
|
|
87
|
-
const
|
|
88
|
-
if (isRecord(
|
|
89
|
-
|
|
90
|
-
if (
|
|
91
|
-
else if (typeof parsed.message === "string") message = parsed.message;
|
|
59
|
+
const value: unknown = JSON.parse(body);
|
|
60
|
+
if (isRecord(value)) {
|
|
61
|
+
if (isRecord(value.error) && typeof value.error.message === "string") message = value.error.message;
|
|
62
|
+
else if (typeof value.message === "string") message = value.message;
|
|
92
63
|
}
|
|
93
64
|
} catch {
|
|
94
65
|
// Plain-text error body.
|
|
@@ -96,70 +67,117 @@ function serverErrorMessage(body: string, token: string): string {
|
|
|
96
67
|
return message.replaceAll(token, "[redacted]").slice(0, MAX_ERROR_MESSAGE_LENGTH).trim();
|
|
97
68
|
}
|
|
98
69
|
|
|
70
|
+
export function detectImageMimeType(bytes: Buffer): GeneratedImage["mimeType"] | undefined {
|
|
71
|
+
if (
|
|
72
|
+
bytes.length >= 8 &&
|
|
73
|
+
bytes.subarray(0, 8).equals(Buffer.from([0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a]))
|
|
74
|
+
) {
|
|
75
|
+
return "image/png";
|
|
76
|
+
}
|
|
77
|
+
if (bytes.length >= 3 && bytes[0] === 0xff && bytes[1] === 0xd8 && bytes[2] === 0xff) return "image/jpeg";
|
|
78
|
+
if (
|
|
79
|
+
bytes.length >= 12 &&
|
|
80
|
+
bytes.subarray(0, 4).toString("ascii") === "RIFF" &&
|
|
81
|
+
bytes.subarray(8, 12).toString("ascii") === "WEBP"
|
|
82
|
+
) {
|
|
83
|
+
return "image/webp";
|
|
84
|
+
}
|
|
85
|
+
return undefined;
|
|
86
|
+
}
|
|
87
|
+
|
|
99
88
|
function decodeImageResponse(value: unknown): GeneratedImage {
|
|
100
|
-
if (!isRecord(value) || !Array.isArray(value.data) ||
|
|
101
|
-
throw new Error("
|
|
89
|
+
if (!isRecord(value) || !Array.isArray(value.data) || !isRecord(value.data[0])) {
|
|
90
|
+
throw new Error("xAI returned an invalid image response");
|
|
102
91
|
}
|
|
103
92
|
const encoded = value.data[0].b64_json;
|
|
104
|
-
if (typeof encoded !== "string" || !encoded.trim())
|
|
105
|
-
throw new Error("OpenAI Codex returned an invalid image response");
|
|
106
|
-
}
|
|
93
|
+
if (typeof encoded !== "string" || !encoded.trim()) throw new Error("xAI returned an invalid image response");
|
|
107
94
|
const base64 = encoded.trim();
|
|
108
95
|
if (!/^[A-Za-z0-9+/]*={0,2}$/.test(base64) || base64.length % 4 === 1) {
|
|
109
|
-
throw new Error("
|
|
96
|
+
throw new Error("xAI returned invalid base64 image data");
|
|
110
97
|
}
|
|
111
98
|
const bytes = Buffer.from(base64, "base64");
|
|
112
99
|
if (bytes.toString("base64").replace(/=+$/, "") !== base64.replace(/=+$/, "")) {
|
|
113
|
-
throw new Error("
|
|
114
|
-
}
|
|
115
|
-
if (bytes.length < PNG_SIGNATURE.length || !bytes.subarray(0, PNG_SIGNATURE.length).equals(PNG_SIGNATURE)) {
|
|
116
|
-
throw new Error("OpenAI Codex returned image data that is not PNG");
|
|
100
|
+
throw new Error("xAI returned invalid base64 image data");
|
|
117
101
|
}
|
|
118
|
-
|
|
102
|
+
const mimeType = detectImageMimeType(bytes);
|
|
103
|
+
if (!mimeType) throw new Error("xAI returned unsupported image data");
|
|
104
|
+
return { bytes, base64: bytes.toString("base64"), mimeType };
|
|
105
|
+
}
|
|
106
|
+
|
|
107
|
+
function retryable(status: number): boolean {
|
|
108
|
+
return status === 408 || status === 409 || status === 425 || status === 429 || status >= 500;
|
|
109
|
+
}
|
|
110
|
+
|
|
111
|
+
async function delay(milliseconds: number, signal?: AbortSignal): Promise<void> {
|
|
112
|
+
await new Promise<void>((resolve, reject) => {
|
|
113
|
+
const timeout = setTimeout(() => {
|
|
114
|
+
signal?.removeEventListener("abort", abort);
|
|
115
|
+
resolve();
|
|
116
|
+
}, milliseconds);
|
|
117
|
+
const abort = () => {
|
|
118
|
+
clearTimeout(timeout);
|
|
119
|
+
signal?.removeEventListener("abort", abort);
|
|
120
|
+
reject(signal?.reason ?? new DOMException("Aborted", "AbortError"));
|
|
121
|
+
};
|
|
122
|
+
if (signal?.aborted) abort();
|
|
123
|
+
else signal?.addEventListener("abort", abort, { once: true });
|
|
124
|
+
});
|
|
119
125
|
}
|
|
120
126
|
|
|
121
127
|
async function requestImage(
|
|
122
128
|
operation: "generation" | "edit",
|
|
123
129
|
body: Record<string, unknown>,
|
|
124
|
-
|
|
125
|
-
signal
|
|
130
|
+
token: string,
|
|
131
|
+
signal?: AbortSignal,
|
|
126
132
|
): Promise<GeneratedImage> {
|
|
127
133
|
const route = operation === "generation" ? "generations" : "edits";
|
|
128
|
-
|
|
129
|
-
|
|
130
|
-
|
|
131
|
-
|
|
132
|
-
|
|
133
|
-
|
|
134
|
-
|
|
135
|
-
|
|
136
|
-
|
|
137
|
-
|
|
138
|
-
|
|
139
|
-
|
|
140
|
-
|
|
141
|
-
|
|
142
|
-
|
|
143
|
-
|
|
144
|
-
|
|
145
|
-
|
|
146
|
-
|
|
147
|
-
|
|
148
|
-
|
|
149
|
-
|
|
134
|
+
for (let attempt = 1; attempt <= MAX_ATTEMPTS; attempt++) {
|
|
135
|
+
const requestSignal = signal
|
|
136
|
+
? AbortSignal.any([signal, AbortSignal.timeout(REQUEST_TIMEOUT_MS)])
|
|
137
|
+
: AbortSignal.timeout(REQUEST_TIMEOUT_MS);
|
|
138
|
+
let response: HttpResponse;
|
|
139
|
+
try {
|
|
140
|
+
response = (await fetch(`${XAI_API_BASE_URL}/images/${route}`, {
|
|
141
|
+
method: "POST",
|
|
142
|
+
headers: {
|
|
143
|
+
Accept: "application/json",
|
|
144
|
+
Authorization: `Bearer ${token}`,
|
|
145
|
+
"Content-Type": "application/json",
|
|
146
|
+
},
|
|
147
|
+
body: JSON.stringify(body),
|
|
148
|
+
signal: requestSignal,
|
|
149
|
+
})) as HttpResponse;
|
|
150
|
+
} catch (error) {
|
|
151
|
+
if (signal?.aborted) throw signal.reason;
|
|
152
|
+
if (attempt === MAX_ATTEMPTS) throw error;
|
|
153
|
+
await delay(RETRY_BASE_DELAY_MS * 2 ** (attempt - 1), signal);
|
|
154
|
+
continue;
|
|
155
|
+
}
|
|
156
|
+
if (response.ok) {
|
|
157
|
+
let value: unknown;
|
|
158
|
+
try {
|
|
159
|
+
value = await response.json();
|
|
160
|
+
} catch {
|
|
161
|
+
throw new Error("xAI returned a non-JSON image response");
|
|
162
|
+
}
|
|
163
|
+
return decodeImageResponse(value);
|
|
164
|
+
}
|
|
165
|
+
const message = serverErrorMessage(await boundedError(response), token);
|
|
166
|
+
if (!retryable(response.status) || attempt === MAX_ATTEMPTS) {
|
|
167
|
+
throw new Error(
|
|
168
|
+
`xAI image ${operation} failed with status ${response.status}${message ? `: ${message}` : ""}`,
|
|
169
|
+
);
|
|
170
|
+
}
|
|
171
|
+
await delay(RETRY_BASE_DELAY_MS * 2 ** (attempt - 1), signal);
|
|
150
172
|
}
|
|
151
|
-
|
|
173
|
+
throw new Error(`xAI image ${operation} failed`);
|
|
152
174
|
}
|
|
153
175
|
|
|
154
|
-
export function generateImage(
|
|
155
|
-
prompt: string,
|
|
156
|
-
auth: CodexAuth,
|
|
157
|
-
signal: AbortSignal | undefined,
|
|
158
|
-
): Promise<GeneratedImage> {
|
|
176
|
+
export function generateImage(prompt: string, token: string, signal?: AbortSignal): Promise<GeneratedImage> {
|
|
159
177
|
return requestImage(
|
|
160
178
|
"generation",
|
|
161
|
-
{
|
|
162
|
-
|
|
179
|
+
{ model: XAI_IMAGE_MODEL, prompt, n: 1, resolution: "1k", response_format: "b64_json" },
|
|
180
|
+
token,
|
|
163
181
|
signal,
|
|
164
182
|
);
|
|
165
183
|
}
|
|
@@ -167,20 +185,21 @@ export function generateImage(
|
|
|
167
185
|
export function editImage(
|
|
168
186
|
prompt: string,
|
|
169
187
|
images: readonly EditImage[],
|
|
170
|
-
|
|
171
|
-
signal
|
|
188
|
+
token: string,
|
|
189
|
+
signal?: AbortSignal,
|
|
172
190
|
): Promise<GeneratedImage> {
|
|
191
|
+
const references = images.map((image) => ({ url: `data:${image.mimeType};base64,${image.data}` }));
|
|
173
192
|
return requestImage(
|
|
174
193
|
"edit",
|
|
175
194
|
{
|
|
176
|
-
|
|
195
|
+
model: XAI_IMAGE_MODEL,
|
|
177
196
|
prompt,
|
|
178
|
-
|
|
179
|
-
|
|
180
|
-
|
|
181
|
-
|
|
197
|
+
n: 1,
|
|
198
|
+
resolution: "1k",
|
|
199
|
+
response_format: "b64_json",
|
|
200
|
+
...(references.length === 1 ? { image: references[0] } : { images: references }),
|
|
182
201
|
},
|
|
183
|
-
|
|
202
|
+
token,
|
|
184
203
|
signal,
|
|
185
204
|
);
|
|
186
205
|
}
|
|
@@ -2,22 +2,21 @@ import { defineTool, withFileMutationQueue, type ExtensionAPI } from "@earendil-
|
|
|
2
2
|
import { randomUUID } from "node:crypto";
|
|
3
3
|
import { link, mkdir, readFile, rm, stat, writeFile } from "node:fs/promises";
|
|
4
4
|
import { homedir } from "node:os";
|
|
5
|
-
import { basename, dirname, isAbsolute, join, resolve } from "node:path";
|
|
5
|
+
import { basename, dirname, extname, isAbsolute, join, resolve } from "node:path";
|
|
6
6
|
import { type Static, Type } from "typebox";
|
|
7
|
-
import {
|
|
7
|
+
import { XAI_IMAGE_MODEL, XAI_PROVIDER } from "../xai/constants.ts";
|
|
8
|
+
import { detectImageMimeType, editImage, generateImage, type EditImage, type GeneratedImage } from "./client.ts";
|
|
8
9
|
|
|
9
|
-
const MODEL = "gpt-image-2";
|
|
10
10
|
const MAX_INPUT_BYTES = 50 * 1024 * 1024;
|
|
11
11
|
const MAX_INLINE_BYTES = 12 * 1024 * 1024;
|
|
12
|
-
const PNG_SIGNATURE = Buffer.from([0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a]);
|
|
13
12
|
|
|
14
13
|
const imageGenSchema = Type.Object(
|
|
15
14
|
{
|
|
16
15
|
prompt: Type.String({ minLength: 1 }),
|
|
17
16
|
path: Type.Optional(
|
|
18
|
-
Type.String({ description: "Explicit
|
|
17
|
+
Type.String({ description: "Explicit image destination path; defaults to Tau's external image store" }),
|
|
19
18
|
),
|
|
20
|
-
referenced_image_paths: Type.Optional(Type.Array(Type.String({ minLength: 1 }), { minItems: 1, maxItems:
|
|
19
|
+
referenced_image_paths: Type.Optional(Type.Array(Type.String({ minLength: 1 }), { minItems: 1, maxItems: 3 })),
|
|
21
20
|
},
|
|
22
21
|
{ additionalProperties: false },
|
|
23
22
|
);
|
|
@@ -26,23 +25,14 @@ type ImageGenParams = Static<typeof imageGenSchema>;
|
|
|
26
25
|
|
|
27
26
|
interface ImageGenDetails {
|
|
28
27
|
path: string;
|
|
29
|
-
model: typeof
|
|
28
|
+
model: typeof XAI_IMAGE_MODEL;
|
|
30
29
|
operation: "generate" | "edit";
|
|
31
30
|
}
|
|
32
31
|
|
|
33
|
-
function
|
|
34
|
-
if (
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
if (bytes.length >= 3 && bytes[0] === 0xff && bytes[1] === 0xd8 && bytes[2] === 0xff) return "image/jpeg";
|
|
38
|
-
if (
|
|
39
|
-
bytes.length >= 12 &&
|
|
40
|
-
bytes.subarray(0, 4).toString("ascii") === "RIFF" &&
|
|
41
|
-
bytes.subarray(8, 12).toString("ascii") === "WEBP"
|
|
42
|
-
) {
|
|
43
|
-
return "image/webp";
|
|
44
|
-
}
|
|
45
|
-
return undefined;
|
|
32
|
+
function outputExtension(image: GeneratedImage): string {
|
|
33
|
+
if (image.mimeType === "image/png") return ".png";
|
|
34
|
+
if (image.mimeType === "image/webp") return ".webp";
|
|
35
|
+
return ".jpg";
|
|
46
36
|
}
|
|
47
37
|
|
|
48
38
|
export default function imageGenExtension(pi: ExtensionAPI): void {
|
|
@@ -51,12 +41,12 @@ export default function imageGenExtension(pi: ExtensionAPI): void {
|
|
|
51
41
|
name: "image_gen",
|
|
52
42
|
label: "Image Generation",
|
|
53
43
|
description:
|
|
54
|
-
"Generate a new raster image from a prompt, or edit one to
|
|
55
|
-
promptSnippet: "Generate or edit raster images with
|
|
44
|
+
"Generate a new raster image from a prompt, or edit one to three local raster images. Uses the xAI Grok subscription OAuth login, saves the result under Tau's external image store unless an explicit path is provided, and returns the image for inspection.",
|
|
45
|
+
promptSnippet: "Generate or edit raster images with Grok Imagine",
|
|
56
46
|
promptGuidelines: [
|
|
57
47
|
"Use image_gen when the user asks for a generated raster image or an AI edit of local raster images.",
|
|
58
48
|
"Omit referenced_image_paths when image_gen should create a new image.",
|
|
59
|
-
"Pass one to
|
|
49
|
+
"Pass one to three local paths in referenced_image_paths when image_gen should edit or compose existing images.",
|
|
60
50
|
"Omit path for temporary external storage. Pass path only when the user wants the generated image saved in their repository or another explicit location.",
|
|
61
51
|
],
|
|
62
52
|
parameters: imageGenSchema,
|
|
@@ -66,18 +56,22 @@ export default function imageGenExtension(pi: ExtensionAPI): void {
|
|
|
66
56
|
if (!prompt) throw new Error("Image prompt cannot be empty");
|
|
67
57
|
const requestedPath = params.path?.startsWith("@") ? params.path.slice(1) : params.path;
|
|
68
58
|
if (requestedPath !== undefined && !requestedPath.trim()) throw new Error("Image path cannot be empty");
|
|
69
|
-
const
|
|
59
|
+
const requestedAbsolutePath = requestedPath
|
|
70
60
|
? isAbsolute(requestedPath)
|
|
71
61
|
? requestedPath
|
|
72
62
|
: resolve(ctx.cwd, requestedPath)
|
|
73
|
-
:
|
|
74
|
-
if (
|
|
63
|
+
: undefined;
|
|
64
|
+
if (
|
|
65
|
+
requestedAbsolutePath &&
|
|
66
|
+
![".jpg", ".jpeg", ".png", ".webp"].includes(extname(requestedAbsolutePath).toLowerCase())
|
|
67
|
+
) {
|
|
68
|
+
throw new Error("Image path must end in .jpg, .jpeg, .png, or .webp");
|
|
69
|
+
}
|
|
75
70
|
|
|
76
|
-
const token = await ctx.modelRegistry.getApiKeyForProvider(
|
|
71
|
+
const token = await ctx.modelRegistry.getApiKeyForProvider(XAI_PROVIDER);
|
|
77
72
|
if (!token) {
|
|
78
|
-
throw new Error("
|
|
73
|
+
throw new Error("xAI OAuth is unavailable. Run /login and select xAI (Grok subscription OAuth).");
|
|
79
74
|
}
|
|
80
|
-
const auth = resolveCodexAuth(token);
|
|
81
75
|
const images: EditImage[] = [];
|
|
82
76
|
for (const path of params.referenced_image_paths ?? []) {
|
|
83
77
|
const rawPath = path.startsWith("@") ? path.slice(1) : path;
|
|
@@ -101,22 +95,37 @@ export default function imageGenExtension(pi: ExtensionAPI): void {
|
|
|
101
95
|
type: "text",
|
|
102
96
|
text:
|
|
103
97
|
operation === "generate"
|
|
104
|
-
? `Generating image with ${
|
|
105
|
-
: `Editing image with ${
|
|
98
|
+
? `Generating image with ${XAI_IMAGE_MODEL}...`
|
|
99
|
+
: `Editing image with ${XAI_IMAGE_MODEL}...`,
|
|
106
100
|
},
|
|
107
101
|
],
|
|
108
102
|
details: undefined,
|
|
109
103
|
});
|
|
110
104
|
const generated =
|
|
111
105
|
operation === "generate"
|
|
112
|
-
? await generateImage(prompt,
|
|
113
|
-
: await editImage(prompt, images,
|
|
106
|
+
? await generateImage(prompt, token, signal)
|
|
107
|
+
: await editImage(prompt, images, token, signal);
|
|
114
108
|
signal?.throwIfAborted();
|
|
109
|
+
const generatedExtension = outputExtension(generated);
|
|
110
|
+
if (requestedAbsolutePath) {
|
|
111
|
+
const requestedExtension = extname(requestedAbsolutePath).toLowerCase();
|
|
112
|
+
const matches =
|
|
113
|
+
requestedExtension === generatedExtension ||
|
|
114
|
+
(generatedExtension === ".jpg" && requestedExtension === ".jpeg");
|
|
115
|
+
if (!matches)
|
|
116
|
+
throw new Error(`xAI returned ${generated.mimeType}; destination must end in ${generatedExtension}`);
|
|
117
|
+
}
|
|
118
|
+
const absolutePath =
|
|
119
|
+
requestedAbsolutePath ??
|
|
120
|
+
join(homedir(), ".local", "share", "tau-agent", "images", `image-${randomUUID()}${generatedExtension}`);
|
|
115
121
|
|
|
116
122
|
const outputDirectory = dirname(absolutePath);
|
|
117
123
|
await withFileMutationQueue(absolutePath, async () => {
|
|
118
124
|
await mkdir(outputDirectory, { recursive: true });
|
|
119
|
-
const temporaryPath = join(
|
|
125
|
+
const temporaryPath = join(
|
|
126
|
+
outputDirectory,
|
|
127
|
+
`.${basename(absolutePath)}.${randomUUID()}.tmp${generatedExtension}`,
|
|
128
|
+
);
|
|
120
129
|
try {
|
|
121
130
|
await writeFile(temporaryPath, generated.bytes, { flag: "wx" });
|
|
122
131
|
signal?.throwIfAborted();
|
|
@@ -127,7 +136,7 @@ export default function imageGenExtension(pi: ExtensionAPI): void {
|
|
|
127
136
|
});
|
|
128
137
|
|
|
129
138
|
const verb = operation === "generate" ? "Generated" : "Edited";
|
|
130
|
-
const details: ImageGenDetails = { path: absolutePath, model:
|
|
139
|
+
const details: ImageGenDetails = { path: absolutePath, model: XAI_IMAGE_MODEL, operation };
|
|
131
140
|
if (generated.bytes.length > MAX_INLINE_BYTES) {
|
|
132
141
|
return {
|
|
133
142
|
content: [
|
|
@@ -46,6 +46,10 @@ Adds `/ideas` to log rough ideas or open the ideas browser.
|
|
|
46
46
|
|
|
47
47
|
Gives the agent an image-generation tool using the configured image service. Generated images are saved for inspection.
|
|
48
48
|
|
|
49
|
+
## xai
|
|
50
|
+
|
|
51
|
+
Adds Grok 4.5 and Grok Imagine through an xAI Grok subscription OAuth login. Run `/login` and select xAI before use.
|
|
52
|
+
|
|
49
53
|
## manage-sessions
|
|
50
54
|
|
|
51
55
|
Adds `/manage-sessions` to browse saved sessions and `/sweep` to archive or delete the current session after starting a new one.
|
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
# xAI OAuth
|
|
2
|
+
|
|
3
|
+
Adds Grok 4.5 using an xAI Grok subscription login. Run `/login`, select **xAI (Grok subscription OAuth)**, then authorize xAI in the browser. An existing official Grok CLI login can also be reused.
|
|
4
|
+
|
|
5
|
+
The same login powers Tau's `image_gen` tool through Grok Imagine. No xAI API key is used.
|
|
6
|
+
|
|
7
|
+
This integration uses xAI's undocumented subscription OAuth access. xAI may change its availability, entitlement rules, or protocol without notice.
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
import type { OAuthCredentials } from "@earendil-works/pi-ai";
|
|
2
|
+
import { readFile } from "node:fs/promises";
|
|
3
|
+
import { homedir } from "node:os";
|
|
4
|
+
import { join } from "node:path";
|
|
5
|
+
import { XAI_OAUTH_CLIENT_ID, XAI_OAUTH_ISSUER } from "./constants.ts";
|
|
6
|
+
|
|
7
|
+
function expiry(value: unknown): number | undefined {
|
|
8
|
+
if (typeof value === "number" && Number.isInteger(value) && value >= 1_000_000_000_000) return value;
|
|
9
|
+
if (typeof value !== "string") return undefined;
|
|
10
|
+
const parsed = Date.parse(value);
|
|
11
|
+
return Number.isFinite(parsed) ? parsed : undefined;
|
|
12
|
+
}
|
|
13
|
+
|
|
14
|
+
export function parseGrokCredentials(value: unknown): OAuthCredentials | undefined {
|
|
15
|
+
if (typeof value !== "object" || value === null || Array.isArray(value)) return undefined;
|
|
16
|
+
const entry = (value as Record<string, unknown>)[`${XAI_OAUTH_ISSUER}::${XAI_OAUTH_CLIENT_ID}`];
|
|
17
|
+
if (typeof entry !== "object" || entry === null || Array.isArray(entry)) return undefined;
|
|
18
|
+
const record = entry as Record<string, unknown>;
|
|
19
|
+
const expires = expiry(record.expires_at);
|
|
20
|
+
if (
|
|
21
|
+
typeof record.key !== "string" ||
|
|
22
|
+
!record.key ||
|
|
23
|
+
typeof record.refresh_token !== "string" ||
|
|
24
|
+
!record.refresh_token ||
|
|
25
|
+
record.oidc_issuer !== XAI_OAUTH_ISSUER ||
|
|
26
|
+
record.oidc_client_id !== XAI_OAUTH_CLIENT_ID ||
|
|
27
|
+
expires === undefined
|
|
28
|
+
) {
|
|
29
|
+
return undefined;
|
|
30
|
+
}
|
|
31
|
+
return { access: record.key, refresh: record.refresh_token, expires };
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
export async function readGrokCredentials(): Promise<OAuthCredentials | undefined> {
|
|
35
|
+
try {
|
|
36
|
+
return parseGrokCredentials(JSON.parse(await readFile(join(homedir(), ".grok", "auth.json"), "utf8")));
|
|
37
|
+
} catch {
|
|
38
|
+
return undefined;
|
|
39
|
+
}
|
|
40
|
+
}
|
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
export const XAI_PROVIDER = "xai-oauth";
|
|
2
|
+
export const XAI_CHAT_MODEL = "grok-4.5";
|
|
3
|
+
export const XAI_IMAGE_MODEL = "grok-imagine-image-quality";
|
|
4
|
+
export const XAI_API_BASE_URL = "https://api.x.ai/v1";
|
|
5
|
+
|
|
6
|
+
export const XAI_OAUTH_ISSUER = "https://auth.x.ai";
|
|
7
|
+
export const XAI_OAUTH_CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828";
|
|
8
|
+
export const XAI_OAUTH_SCOPE = "openid profile email offline_access grok-cli:access api:access";
|
|
9
|
+
export const XAI_OAUTH_CALLBACK_HOST = "127.0.0.1";
|
|
10
|
+
export const XAI_OAUTH_CALLBACK_PORT = 56121;
|
|
11
|
+
export const XAI_OAUTH_CALLBACK_PATH = "/callback";
|
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
import type { ExtensionAPI } from "@earendil-works/pi-coding-agent";
|
|
2
|
+
import { XAI_API_BASE_URL, XAI_CHAT_MODEL, XAI_PROVIDER } from "./constants.ts";
|
|
3
|
+
import { xaiOAuth } from "./oauth.ts";
|
|
4
|
+
import { rewriteXaiPayload } from "./payload.ts";
|
|
5
|
+
|
|
6
|
+
export default function xaiExtension(pi: ExtensionAPI): void {
|
|
7
|
+
pi.registerProvider(XAI_PROVIDER, {
|
|
8
|
+
name: "xAI (Grok subscription OAuth)",
|
|
9
|
+
baseUrl: XAI_API_BASE_URL,
|
|
10
|
+
api: "openai-responses",
|
|
11
|
+
authHeader: true,
|
|
12
|
+
oauth: xaiOAuth,
|
|
13
|
+
models: [
|
|
14
|
+
{
|
|
15
|
+
id: XAI_CHAT_MODEL,
|
|
16
|
+
name: "Grok 4.5",
|
|
17
|
+
reasoning: true,
|
|
18
|
+
input: ["text", "image"],
|
|
19
|
+
cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 },
|
|
20
|
+
contextWindow: 500_000,
|
|
21
|
+
maxTokens: 131_072,
|
|
22
|
+
thinkingLevelMap: {
|
|
23
|
+
off: null,
|
|
24
|
+
minimal: "low",
|
|
25
|
+
low: "low",
|
|
26
|
+
medium: "medium",
|
|
27
|
+
high: "high",
|
|
28
|
+
xhigh: null,
|
|
29
|
+
max: null,
|
|
30
|
+
},
|
|
31
|
+
},
|
|
32
|
+
],
|
|
33
|
+
});
|
|
34
|
+
pi.on("before_provider_request", (event, ctx) => {
|
|
35
|
+
if (ctx.model?.provider !== XAI_PROVIDER) return;
|
|
36
|
+
return rewriteXaiPayload(event.payload);
|
|
37
|
+
});
|
|
38
|
+
}
|
|
@@ -0,0 +1,342 @@
|
|
|
1
|
+
import type { OAuthCredentials, OAuthLoginCallbacks } from "@earendil-works/pi-ai";
|
|
2
|
+
import { createHash, randomBytes } from "node:crypto";
|
|
3
|
+
import { createServer, type Server } from "node:http";
|
|
4
|
+
import { readGrokCredentials } from "./auth.ts";
|
|
5
|
+
import {
|
|
6
|
+
XAI_OAUTH_CALLBACK_HOST,
|
|
7
|
+
XAI_OAUTH_CALLBACK_PATH,
|
|
8
|
+
XAI_OAUTH_CALLBACK_PORT,
|
|
9
|
+
XAI_OAUTH_CLIENT_ID,
|
|
10
|
+
XAI_OAUTH_ISSUER,
|
|
11
|
+
XAI_OAUTH_SCOPE,
|
|
12
|
+
} from "./constants.ts";
|
|
13
|
+
|
|
14
|
+
const DISCOVERY_URL = `${XAI_OAUTH_ISSUER}/.well-known/openid-configuration`;
|
|
15
|
+
const REQUEST_TIMEOUT_MS = 30_000;
|
|
16
|
+
const LOGIN_TIMEOUT_MS = 180_000;
|
|
17
|
+
const REFRESH_SKEW_MS = 120_000;
|
|
18
|
+
|
|
19
|
+
interface Discovery {
|
|
20
|
+
authorization_endpoint: string;
|
|
21
|
+
token_endpoint: string;
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
interface TokenPayload {
|
|
25
|
+
access_token?: unknown;
|
|
26
|
+
refresh_token?: unknown;
|
|
27
|
+
id_token?: unknown;
|
|
28
|
+
expires_in?: unknown;
|
|
29
|
+
token_type?: unknown;
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
interface HttpResponse {
|
|
33
|
+
ok: boolean;
|
|
34
|
+
status: number;
|
|
35
|
+
body: { cancel(): Promise<void> } | null;
|
|
36
|
+
json(): Promise<unknown>;
|
|
37
|
+
}
|
|
38
|
+
|
|
39
|
+
interface CallbackResult {
|
|
40
|
+
code?: string;
|
|
41
|
+
error?: string;
|
|
42
|
+
errorDescription?: string;
|
|
43
|
+
}
|
|
44
|
+
|
|
45
|
+
function validatedEndpoint(value: unknown, field: string): string {
|
|
46
|
+
if (typeof value !== "string") throw new Error(`xAI OAuth discovery omitted ${field}`);
|
|
47
|
+
const url = new URL(value);
|
|
48
|
+
const host = url.hostname.toLowerCase();
|
|
49
|
+
if (url.protocol !== "https:" || (host !== "x.ai" && !host.endsWith(".x.ai"))) {
|
|
50
|
+
throw new Error(`xAI OAuth discovery returned an unexpected ${field}`);
|
|
51
|
+
}
|
|
52
|
+
return url.toString();
|
|
53
|
+
}
|
|
54
|
+
|
|
55
|
+
function requestSignal(parent?: AbortSignal): AbortSignal {
|
|
56
|
+
const timeout = AbortSignal.timeout(REQUEST_TIMEOUT_MS);
|
|
57
|
+
return parent ? AbortSignal.any([parent, timeout]) : timeout;
|
|
58
|
+
}
|
|
59
|
+
|
|
60
|
+
async function discover(signal?: AbortSignal): Promise<Discovery> {
|
|
61
|
+
const response = (await fetch(DISCOVERY_URL, {
|
|
62
|
+
headers: { Accept: "application/json" },
|
|
63
|
+
signal: requestSignal(signal),
|
|
64
|
+
})) as HttpResponse;
|
|
65
|
+
if (!response.ok) throw new Error(`xAI OAuth discovery failed with status ${response.status}`);
|
|
66
|
+
const value: unknown = await response.json();
|
|
67
|
+
if (typeof value !== "object" || value === null || Array.isArray(value)) {
|
|
68
|
+
throw new Error("xAI OAuth discovery returned invalid JSON");
|
|
69
|
+
}
|
|
70
|
+
const record = value as Record<string, unknown>;
|
|
71
|
+
return {
|
|
72
|
+
authorization_endpoint: validatedEndpoint(record.authorization_endpoint, "authorization_endpoint"),
|
|
73
|
+
token_endpoint: validatedEndpoint(record.token_endpoint, "token_endpoint"),
|
|
74
|
+
};
|
|
75
|
+
}
|
|
76
|
+
|
|
77
|
+
async function tokenRequest(endpoint: string, body: URLSearchParams, signal?: AbortSignal): Promise<TokenPayload> {
|
|
78
|
+
const response = (await fetch(validatedEndpoint(endpoint, "token_endpoint"), {
|
|
79
|
+
method: "POST",
|
|
80
|
+
headers: { Accept: "application/json", "Content-Type": "application/x-www-form-urlencoded" },
|
|
81
|
+
body,
|
|
82
|
+
signal: requestSignal(signal),
|
|
83
|
+
})) as HttpResponse;
|
|
84
|
+
if (!response.ok) {
|
|
85
|
+
await response.body?.cancel().catch(() => undefined);
|
|
86
|
+
throw new Error(`xAI OAuth token request failed with status ${response.status}`);
|
|
87
|
+
}
|
|
88
|
+
return (await response.json()) as TokenPayload;
|
|
89
|
+
}
|
|
90
|
+
|
|
91
|
+
function jwtClaims(token: string): Record<string, unknown> {
|
|
92
|
+
const segments = token.split(".");
|
|
93
|
+
if (segments.length !== 3 || !segments[1]) throw new Error("xAI OAuth returned an invalid ID token");
|
|
94
|
+
try {
|
|
95
|
+
const value: unknown = JSON.parse(Buffer.from(segments[1], "base64url").toString("utf8"));
|
|
96
|
+
if (typeof value !== "object" || value === null || Array.isArray(value)) throw new Error();
|
|
97
|
+
return value as Record<string, unknown>;
|
|
98
|
+
} catch {
|
|
99
|
+
throw new Error("xAI OAuth returned an invalid ID token");
|
|
100
|
+
}
|
|
101
|
+
}
|
|
102
|
+
|
|
103
|
+
function credentials(payload: TokenPayload, endpoint: string, fallbackRefresh = "", nonce?: string): OAuthCredentials {
|
|
104
|
+
if (typeof payload.access_token !== "string" || !payload.access_token) {
|
|
105
|
+
throw new Error("xAI OAuth token response omitted the access token");
|
|
106
|
+
}
|
|
107
|
+
const refresh =
|
|
108
|
+
typeof payload.refresh_token === "string" && payload.refresh_token ? payload.refresh_token : fallbackRefresh;
|
|
109
|
+
if (!refresh) throw new Error("xAI OAuth token response omitted the refresh token");
|
|
110
|
+
if (nonce !== undefined) {
|
|
111
|
+
if (typeof payload.id_token !== "string" || !payload.id_token)
|
|
112
|
+
throw new Error("xAI OAuth token response omitted the ID token");
|
|
113
|
+
const claims = jwtClaims(payload.id_token);
|
|
114
|
+
const audience = claims.aud;
|
|
115
|
+
const validAudience =
|
|
116
|
+
audience === XAI_OAUTH_CLIENT_ID || (Array.isArray(audience) && audience.includes(XAI_OAUTH_CLIENT_ID));
|
|
117
|
+
if (claims.iss !== XAI_OAUTH_ISSUER || !validAudience || claims.nonce !== nonce) {
|
|
118
|
+
throw new Error("xAI OAuth ID token validation failed");
|
|
119
|
+
}
|
|
120
|
+
if (typeof claims.exp !== "number" || claims.exp * 1000 <= Date.now()) {
|
|
121
|
+
throw new Error("xAI OAuth returned an expired ID token");
|
|
122
|
+
}
|
|
123
|
+
}
|
|
124
|
+
const expiresIn = typeof payload.expires_in === "number" && payload.expires_in > 0 ? payload.expires_in : 3600;
|
|
125
|
+
return {
|
|
126
|
+
access: payload.access_token,
|
|
127
|
+
refresh,
|
|
128
|
+
expires: Date.now() + expiresIn * 1000 - REFRESH_SKEW_MS,
|
|
129
|
+
tokenEndpoint: endpoint,
|
|
130
|
+
};
|
|
131
|
+
}
|
|
132
|
+
|
|
133
|
+
async function closeServer(server: Server): Promise<void> {
|
|
134
|
+
if (!server.listening) return;
|
|
135
|
+
await new Promise<void>((resolve) => server.close(() => resolve()));
|
|
136
|
+
}
|
|
137
|
+
|
|
138
|
+
async function callbackServer(expectedState: string): Promise<{
|
|
139
|
+
redirectUri: string;
|
|
140
|
+
wait(signal?: AbortSignal): Promise<CallbackResult>;
|
|
141
|
+
acceptManual(input: string): string | undefined;
|
|
142
|
+
close(): Promise<void>;
|
|
143
|
+
}> {
|
|
144
|
+
let settle: ((result: CallbackResult) => void) | undefined;
|
|
145
|
+
let reject: ((error: Error) => void) | undefined;
|
|
146
|
+
let settled = false;
|
|
147
|
+
const result = new Promise<CallbackResult>((resolve, rejectResult) => {
|
|
148
|
+
settle = resolve;
|
|
149
|
+
reject = rejectResult;
|
|
150
|
+
});
|
|
151
|
+
const accept = (value: CallbackResult) => {
|
|
152
|
+
if (settled) return;
|
|
153
|
+
settled = true;
|
|
154
|
+
settle?.(value);
|
|
155
|
+
};
|
|
156
|
+
const parse = (params: URLSearchParams): CallbackResult | undefined => {
|
|
157
|
+
if (params.get("state") !== expectedState) return undefined;
|
|
158
|
+
const code = params.get("code") || undefined;
|
|
159
|
+
const error = params.get("error") || undefined;
|
|
160
|
+
if (!code && !error) return undefined;
|
|
161
|
+
return { code, error, errorDescription: params.get("error_description") || undefined };
|
|
162
|
+
};
|
|
163
|
+
const server = createServer((request, response) => {
|
|
164
|
+
const origin = request.headers.origin;
|
|
165
|
+
if (origin === "https://accounts.x.ai" || origin === "https://auth.x.ai") {
|
|
166
|
+
response.setHeader("Access-Control-Allow-Origin", origin);
|
|
167
|
+
response.setHeader("Access-Control-Allow-Methods", "GET, OPTIONS");
|
|
168
|
+
response.setHeader("Access-Control-Allow-Headers", "Content-Type");
|
|
169
|
+
response.setHeader("Access-Control-Allow-Private-Network", "true");
|
|
170
|
+
response.setHeader("Vary", "Origin");
|
|
171
|
+
}
|
|
172
|
+
if (request.method === "OPTIONS") {
|
|
173
|
+
response.writeHead(204).end();
|
|
174
|
+
return;
|
|
175
|
+
}
|
|
176
|
+
const url = new URL(request.url ?? "/", `http://${XAI_OAUTH_CALLBACK_HOST}`);
|
|
177
|
+
if (request.method !== "GET" || url.pathname !== XAI_OAUTH_CALLBACK_PATH) {
|
|
178
|
+
response.writeHead(404).end("Not found");
|
|
179
|
+
return;
|
|
180
|
+
}
|
|
181
|
+
const parsed = parse(url.searchParams);
|
|
182
|
+
if (!parsed) {
|
|
183
|
+
response.writeHead(400, { "Content-Type": "text/plain; charset=utf-8" }).end("Invalid OAuth callback");
|
|
184
|
+
return;
|
|
185
|
+
}
|
|
186
|
+
response
|
|
187
|
+
.writeHead(parsed.error ? 400 : 200, { "Content-Type": "text/html; charset=utf-8" })
|
|
188
|
+
.end("<html><body><h1>xAI authorization received.</h1>You can close this tab.</body></html>", () =>
|
|
189
|
+
accept(parsed),
|
|
190
|
+
);
|
|
191
|
+
});
|
|
192
|
+
const listen = (port: number) =>
|
|
193
|
+
new Promise<number>((resolve, rejectListen) => {
|
|
194
|
+
server.once("error", rejectListen);
|
|
195
|
+
server.listen(port, XAI_OAUTH_CALLBACK_HOST, () => {
|
|
196
|
+
server.removeListener("error", rejectListen);
|
|
197
|
+
const address = server.address();
|
|
198
|
+
if (!address || typeof address === "string") rejectListen(new Error("Could not determine callback port"));
|
|
199
|
+
else resolve(address.port);
|
|
200
|
+
});
|
|
201
|
+
});
|
|
202
|
+
let port: number;
|
|
203
|
+
try {
|
|
204
|
+
port = await listen(XAI_OAUTH_CALLBACK_PORT);
|
|
205
|
+
} catch {
|
|
206
|
+
port = await listen(0);
|
|
207
|
+
}
|
|
208
|
+
return {
|
|
209
|
+
redirectUri: `http://${XAI_OAUTH_CALLBACK_HOST}:${port}${XAI_OAUTH_CALLBACK_PATH}`,
|
|
210
|
+
acceptManual(input) {
|
|
211
|
+
try {
|
|
212
|
+
const value = input.trim();
|
|
213
|
+
const url = value.startsWith("http")
|
|
214
|
+
? new URL(value)
|
|
215
|
+
: new URL(`http://${XAI_OAUTH_CALLBACK_HOST}${XAI_OAUTH_CALLBACK_PATH}?${value.replace(/^\?/, "")}`);
|
|
216
|
+
if (url.pathname !== XAI_OAUTH_CALLBACK_PATH) return "Callback URL path was not recognized";
|
|
217
|
+
const parsed = parse(url.searchParams);
|
|
218
|
+
if (!parsed) return "Callback state did not match";
|
|
219
|
+
accept(parsed);
|
|
220
|
+
return undefined;
|
|
221
|
+
} catch {
|
|
222
|
+
return "Callback URL was invalid";
|
|
223
|
+
}
|
|
224
|
+
},
|
|
225
|
+
async wait(signal) {
|
|
226
|
+
const timeout = setTimeout(() => {
|
|
227
|
+
if (!settled) {
|
|
228
|
+
settled = true;
|
|
229
|
+
reject?.(new Error("Timed out waiting for xAI OAuth callback"));
|
|
230
|
+
}
|
|
231
|
+
}, LOGIN_TIMEOUT_MS);
|
|
232
|
+
const onAbort = () => {
|
|
233
|
+
if (!settled) {
|
|
234
|
+
settled = true;
|
|
235
|
+
reject?.(new Error("xAI OAuth login was cancelled"));
|
|
236
|
+
}
|
|
237
|
+
};
|
|
238
|
+
signal?.addEventListener("abort", onAbort, { once: true });
|
|
239
|
+
try {
|
|
240
|
+
return await result;
|
|
241
|
+
} finally {
|
|
242
|
+
clearTimeout(timeout);
|
|
243
|
+
signal?.removeEventListener("abort", onAbort);
|
|
244
|
+
await closeServer(server);
|
|
245
|
+
}
|
|
246
|
+
},
|
|
247
|
+
close: () => closeServer(server),
|
|
248
|
+
};
|
|
249
|
+
}
|
|
250
|
+
|
|
251
|
+
async function refreshXaiCredentials(value: OAuthCredentials): Promise<OAuthCredentials> {
|
|
252
|
+
if (!value.refresh) throw new Error("xAI OAuth credential cannot be refreshed; run /login again");
|
|
253
|
+
const endpoint =
|
|
254
|
+
typeof value.tokenEndpoint === "string" && value.tokenEndpoint
|
|
255
|
+
? validatedEndpoint(value.tokenEndpoint, "token_endpoint")
|
|
256
|
+
: (await discover()).token_endpoint;
|
|
257
|
+
const payload = await tokenRequest(
|
|
258
|
+
endpoint,
|
|
259
|
+
new URLSearchParams({
|
|
260
|
+
grant_type: "refresh_token",
|
|
261
|
+
client_id: XAI_OAUTH_CLIENT_ID,
|
|
262
|
+
refresh_token: value.refresh,
|
|
263
|
+
}),
|
|
264
|
+
);
|
|
265
|
+
return credentials(payload, endpoint, value.refresh);
|
|
266
|
+
}
|
|
267
|
+
|
|
268
|
+
export const xaiOAuth = {
|
|
269
|
+
name: "xAI (Grok subscription)",
|
|
270
|
+
usesCallbackServer: true,
|
|
271
|
+
async login(callbacks: OAuthLoginCallbacks): Promise<OAuthCredentials> {
|
|
272
|
+
const existing = await readGrokCredentials();
|
|
273
|
+
if (existing) {
|
|
274
|
+
const method = await callbacks.onSelect({
|
|
275
|
+
message: "Select xAI login method:",
|
|
276
|
+
options: [
|
|
277
|
+
{ id: "browser", label: "Browser login" },
|
|
278
|
+
{ id: "existing", label: "Use existing Grok CLI login" },
|
|
279
|
+
],
|
|
280
|
+
});
|
|
281
|
+
if (!method) throw new Error("Login cancelled");
|
|
282
|
+
if (method === "existing") {
|
|
283
|
+
if (existing.expires > Date.now()) return existing;
|
|
284
|
+
try {
|
|
285
|
+
return await refreshXaiCredentials(existing);
|
|
286
|
+
} catch {
|
|
287
|
+
callbacks.onProgress?.("The existing Grok CLI login could not be refreshed. Starting browser login.");
|
|
288
|
+
}
|
|
289
|
+
}
|
|
290
|
+
}
|
|
291
|
+
const discovery = await discover(callbacks.signal);
|
|
292
|
+
const verifier = randomBytes(32).toString("base64url");
|
|
293
|
+
const challenge = createHash("sha256").update(verifier).digest("base64url");
|
|
294
|
+
const state = randomBytes(24).toString("base64url");
|
|
295
|
+
const nonce = randomBytes(24).toString("base64url");
|
|
296
|
+
const callback = await callbackServer(state);
|
|
297
|
+
try {
|
|
298
|
+
const url = new URL(discovery.authorization_endpoint);
|
|
299
|
+
url.search = new URLSearchParams({
|
|
300
|
+
response_type: "code",
|
|
301
|
+
client_id: XAI_OAUTH_CLIENT_ID,
|
|
302
|
+
redirect_uri: callback.redirectUri,
|
|
303
|
+
scope: XAI_OAUTH_SCOPE,
|
|
304
|
+
code_challenge: challenge,
|
|
305
|
+
code_challenge_method: "S256",
|
|
306
|
+
state,
|
|
307
|
+
nonce,
|
|
308
|
+
}).toString();
|
|
309
|
+
callbacks.onAuth({ url: url.toString(), instructions: "Authorize xAI in your browser, then return to Tau." });
|
|
310
|
+
if (callbacks.onManualCodeInput) {
|
|
311
|
+
void callbacks
|
|
312
|
+
.onManualCodeInput()
|
|
313
|
+
.then((input) => {
|
|
314
|
+
const error = callback.acceptManual(input);
|
|
315
|
+
if (error) callbacks.onProgress?.(`Ignored pasted callback: ${error}`);
|
|
316
|
+
})
|
|
317
|
+
.catch(() => undefined);
|
|
318
|
+
}
|
|
319
|
+
const result = await callback.wait(callbacks.signal);
|
|
320
|
+
if (result.error) throw new Error(`xAI authorization failed: ${result.errorDescription ?? result.error}`);
|
|
321
|
+
if (!result.code) throw new Error("xAI authorization did not return a code");
|
|
322
|
+
const payload = await tokenRequest(
|
|
323
|
+
discovery.token_endpoint,
|
|
324
|
+
new URLSearchParams({
|
|
325
|
+
grant_type: "authorization_code",
|
|
326
|
+
client_id: XAI_OAUTH_CLIENT_ID,
|
|
327
|
+
code: result.code,
|
|
328
|
+
redirect_uri: callback.redirectUri,
|
|
329
|
+
code_verifier: verifier,
|
|
330
|
+
}),
|
|
331
|
+
callbacks.signal,
|
|
332
|
+
);
|
|
333
|
+
return credentials(payload, discovery.token_endpoint, "", nonce);
|
|
334
|
+
} finally {
|
|
335
|
+
await callback.close();
|
|
336
|
+
}
|
|
337
|
+
},
|
|
338
|
+
refreshToken: refreshXaiCredentials,
|
|
339
|
+
getApiKey(value: OAuthCredentials): string {
|
|
340
|
+
return value.access;
|
|
341
|
+
},
|
|
342
|
+
};
|
|
@@ -0,0 +1,68 @@
|
|
|
1
|
+
function isRecord(value: unknown): value is Record<string, unknown> {
|
|
2
|
+
return typeof value === "object" && value !== null && !Array.isArray(value);
|
|
3
|
+
}
|
|
4
|
+
|
|
5
|
+
function contentText(value: unknown): string {
|
|
6
|
+
if (typeof value === "string") return value;
|
|
7
|
+
if (!Array.isArray(value)) return "";
|
|
8
|
+
return value
|
|
9
|
+
.map((part) => {
|
|
10
|
+
if (!isRecord(part)) return "";
|
|
11
|
+
return typeof part.text === "string" ? part.text : "";
|
|
12
|
+
})
|
|
13
|
+
.filter(Boolean)
|
|
14
|
+
.join("\n");
|
|
15
|
+
}
|
|
16
|
+
|
|
17
|
+
function normalizeToolOutput(item: Record<string, unknown>): unknown[] {
|
|
18
|
+
if (item.type !== "function_call_output" || !Array.isArray(item.output)) return [item];
|
|
19
|
+
const images = item.output.filter((part) => isRecord(part) && part.type === "input_image");
|
|
20
|
+
if (images.length === 0) return [item];
|
|
21
|
+
const text = contentText(item.output) || "(tool returned image output)";
|
|
22
|
+
return [
|
|
23
|
+
{ ...item, output: text },
|
|
24
|
+
{
|
|
25
|
+
role: "user",
|
|
26
|
+
content: [
|
|
27
|
+
{ type: "input_text", text: "The previous tool result included image output. Use the attached image." },
|
|
28
|
+
...images,
|
|
29
|
+
],
|
|
30
|
+
},
|
|
31
|
+
];
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
export function rewriteXaiPayload(value: unknown): unknown {
|
|
35
|
+
if (!isRecord(value)) return value;
|
|
36
|
+
const body = { ...value };
|
|
37
|
+
delete body.prompt_cache_retention;
|
|
38
|
+
if (isRecord(body.reasoning)) {
|
|
39
|
+
const effort = body.reasoning.effort;
|
|
40
|
+
body.reasoning =
|
|
41
|
+
typeof effort === "string" && effort !== "none"
|
|
42
|
+
? { effort: effort === "minimal" ? "low" : effort }
|
|
43
|
+
: undefined;
|
|
44
|
+
}
|
|
45
|
+
if (Array.isArray(body.input)) {
|
|
46
|
+
const instructions: string[] = [];
|
|
47
|
+
const input: unknown[] = [];
|
|
48
|
+
for (const raw of body.input) {
|
|
49
|
+
if (!isRecord(raw)) {
|
|
50
|
+
input.push(raw);
|
|
51
|
+
continue;
|
|
52
|
+
}
|
|
53
|
+
if ((raw.role === "developer" || raw.role === "system") && input.length === 0) {
|
|
54
|
+
const text = contentText(raw.content).trim();
|
|
55
|
+
if (text) instructions.push(text);
|
|
56
|
+
continue;
|
|
57
|
+
}
|
|
58
|
+
input.push(...normalizeToolOutput(raw));
|
|
59
|
+
}
|
|
60
|
+
body.input = input;
|
|
61
|
+
if (instructions.length > 0) {
|
|
62
|
+
body.instructions = [typeof body.instructions === "string" ? body.instructions : "", ...instructions]
|
|
63
|
+
.filter(Boolean)
|
|
64
|
+
.join("\n\n");
|
|
65
|
+
}
|
|
66
|
+
}
|
|
67
|
+
return body;
|
|
68
|
+
}
|
package/package.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
{
|
|
2
2
|
"name": "@shanepadgett/tau-agent",
|
|
3
|
-
"version": "0.
|
|
3
|
+
"version": "0.7.0",
|
|
4
4
|
"description": "Tau is a custom agentic harness built with pi extensions",
|
|
5
5
|
"type": "module",
|
|
6
6
|
"license": "MIT",
|
|
@@ -28,7 +28,7 @@
|
|
|
28
28
|
"README.md"
|
|
29
29
|
],
|
|
30
30
|
"dependencies": {
|
|
31
|
-
"@shanepadgett/tau-tui": "0.
|
|
31
|
+
"@shanepadgett/tau-tui": "0.7.0",
|
|
32
32
|
"@toon-format/toon": "2.3.0",
|
|
33
33
|
"smol-toml": "1.7.0"
|
|
34
34
|
},
|
|
@@ -10,21 +10,23 @@ const MAX_ATTEMPTS = 5;
|
|
|
10
10
|
const MAX_TOOL_ATTEMPTS = 2;
|
|
11
11
|
const SEVEN_DAYS_MS = 604_800_000;
|
|
12
12
|
|
|
13
|
-
const PREFERRED_MODELS: ReadonlyArray<{ provider: string; model: string; reasoning: ThinkingLevel }> = [
|
|
14
|
-
{ provider: "openrouter", model: "cohere/north-mini-code:free", reasoning: "high" },
|
|
15
|
-
{ provider: "github-copilot", model: "gemini-3.5-flash", reasoning: "high" },
|
|
16
|
-
{ provider: "openai-codex", model: "gpt-5.4-mini", reasoning: "high" },
|
|
17
|
-
{ provider: "anthropic", model: "claude-haiku-4-5", reasoning: "high" },
|
|
18
|
-
];
|
|
19
|
-
|
|
20
13
|
interface GenerationContext {
|
|
21
14
|
ui: ExtensionContext["ui"];
|
|
22
15
|
signal: AbortSignal | undefined;
|
|
16
|
+
sessionManager?: { getSessionId(): string };
|
|
17
|
+
}
|
|
18
|
+
|
|
19
|
+
interface ModelFallbackOptions {
|
|
20
|
+
statusKey?: string;
|
|
21
|
+
notifyOnFallback?: boolean;
|
|
22
|
+
maxAttempts?: number;
|
|
23
|
+
onStatus?: (status: string) => void | Promise<void>;
|
|
23
24
|
}
|
|
24
25
|
|
|
25
26
|
export async function resolveCandidates(
|
|
26
27
|
ctx: Pick<ExtensionContext, "modelRegistry" | "model" | "cwd" | "isProjectTrusted">,
|
|
27
|
-
preferredModels: ReadonlyArray<{ provider: string; model: string; reasoning: ThinkingLevel }
|
|
28
|
+
preferredModels: ReadonlyArray<{ provider: string; model: string; reasoning: ThinkingLevel }>,
|
|
29
|
+
includeParentModel: boolean,
|
|
28
30
|
): Promise<ModelCandidate[]> {
|
|
29
31
|
const settings = await loadTauExtensionSettings(ctx, modelFallbackSettings);
|
|
30
32
|
const candidates: ModelCandidate[] = [];
|
|
@@ -47,7 +49,7 @@ export async function resolveCandidates(
|
|
|
47
49
|
const model = ctx.modelRegistry.find(preferred.provider, preferred.model);
|
|
48
50
|
if (model) await add(model, preferred.reasoning);
|
|
49
51
|
}
|
|
50
|
-
if (ctx.model) await add(ctx.model, undefined);
|
|
52
|
+
if (includeParentModel && ctx.model) await add(ctx.model, undefined);
|
|
51
53
|
|
|
52
54
|
if (candidates.length === 0) throw new Error("No authenticated model available for generation.");
|
|
53
55
|
return candidates;
|
|
@@ -59,7 +61,7 @@ export async function generateValidated<T>(
|
|
|
59
61
|
prompt: string,
|
|
60
62
|
validate: (text: string) => T,
|
|
61
63
|
correctionPrompt?: (error: Error, text: string) => string,
|
|
62
|
-
options?:
|
|
64
|
+
options?: ModelFallbackOptions,
|
|
63
65
|
): Promise<T> {
|
|
64
66
|
return withModelFallback(ctx, candidates, options, (candidate) =>
|
|
65
67
|
requestValidated(ctx, candidate, prompt, validate, correctionPrompt),
|
|
@@ -73,7 +75,7 @@ export async function generateToolValidated<T>(
|
|
|
73
75
|
tool: Tool,
|
|
74
76
|
validate: (input: unknown) => T,
|
|
75
77
|
correctionPrompt?: (error: Error, output: string) => string,
|
|
76
|
-
options?:
|
|
78
|
+
options?: ModelFallbackOptions,
|
|
77
79
|
): Promise<T> {
|
|
78
80
|
return withModelFallback(ctx, candidates, options, (candidate) =>
|
|
79
81
|
requestToolValidated(
|
|
@@ -91,7 +93,7 @@ export async function generateToolValidated<T>(
|
|
|
91
93
|
async function withModelFallback<T>(
|
|
92
94
|
ctx: GenerationContext,
|
|
93
95
|
candidates: readonly ModelCandidate[],
|
|
94
|
-
options:
|
|
96
|
+
options: ModelFallbackOptions | undefined,
|
|
95
97
|
request: (candidate: ModelCandidate) => Promise<T>,
|
|
96
98
|
): Promise<T> {
|
|
97
99
|
const failures: string[] = [];
|
|
@@ -101,6 +103,7 @@ async function withModelFallback<T>(
|
|
|
101
103
|
for (const [index, candidate] of candidates.entries()) {
|
|
102
104
|
const label = `${candidate.model.provider}/${candidate.model.id}`;
|
|
103
105
|
if (statusKey) ctx.ui.setStatus(statusKey, `generating (${label})`);
|
|
106
|
+
await options?.onStatus?.(`Generating with ${label}`);
|
|
104
107
|
|
|
105
108
|
try {
|
|
106
109
|
return await request(candidate);
|
|
@@ -109,6 +112,7 @@ async function withModelFallback<T>(
|
|
|
109
112
|
if (shouldCooldownProvider(error)) await markProviderUnavailable(candidate.model.provider);
|
|
110
113
|
const message = errorText(error);
|
|
111
114
|
failures.push(`- ${label}: ${message}`);
|
|
115
|
+
if (index < candidates.length - 1) await options?.onStatus?.(`Model failed (${label}); trying next model`);
|
|
112
116
|
if (index < candidates.length - 1 && notifyOnFallback) {
|
|
113
117
|
ctx.ui.notify(`Model failed (${label}): ${message}\nTrying next model.`, "info");
|
|
114
118
|
}
|
|
@@ -209,6 +213,7 @@ function completeCandidate(
|
|
|
209
213
|
headers: candidate.headers,
|
|
210
214
|
signal: ctx.signal,
|
|
211
215
|
reasoning: candidate.reasoning,
|
|
216
|
+
sessionId: ctx.sessionManager?.getSessionId(),
|
|
212
217
|
});
|
|
213
218
|
}
|
|
214
219
|
|