feat: update codex scripts to write model_provider at top level in config files
This commit is contained in:
1 parent
dc4fe57c99
commit
2f6ec66940
3 files changed
+150
-35
No files matched your search
@@ -30,34 +30,81 @@ if ($null -eq $content) {
|
||||
$content = ""
|
||||
}
|
||||
|
||||
if ($content -match '(?m)^[ \t]*model_provider[ \t]*=') {
|
||||
$content = [regex]::Replace(
|
||||
$content,
|
||||
'(?m)^[ \t]*model_provider[ \t]*=.*$',
|
||||
'model_provider = "xixu"'
|
||||
)
|
||||
} else {
|
||||
if ($content.Length -gt 0 -and -not $content.EndsWith("`n")) {
|
||||
$content += "`r`n"
|
||||
$content = $content -replace "`r`n", "`n"
|
||||
$lines = [System.Collections.Generic.List[string]]::new()
|
||||
if ($content.Length -gt 0) {
|
||||
foreach ($line in ($content -split "`n")) {
|
||||
$lines.Add($line)
|
||||
}
|
||||
$content += 'model_provider = "xixu"' + "`r`n"
|
||||
}
|
||||
|
||||
$content = [regex]::Replace(
|
||||
$content,
|
||||
'(?ms)^\[model_providers\.xixu\]\s*.*?(?=^\[|\z)',
|
||||
''
|
||||
)
|
||||
$filteredLines = [System.Collections.Generic.List[string]]::new()
|
||||
$inXixuBlock = $false
|
||||
|
||||
$xixuBlock = @'
|
||||
foreach ($line in $lines) {
|
||||
if ($inXixuBlock) {
|
||||
if ($line -match '^\[.*\][ \t]*$') {
|
||||
$inXixuBlock = $false
|
||||
} else {
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
[model_providers.xixu]
|
||||
name = "Xi Xu's AI Inference"
|
||||
base_url = "https://api.xi-xu.me/v1"
|
||||
env_key = "XIXU_API_KEY"
|
||||
'@
|
||||
if ($line -match '^[ \t]*model_provider[ \t]*=') {
|
||||
continue
|
||||
}
|
||||
|
||||
$content = $content.TrimEnd() + "`r`n" + $xixuBlock + "`r`n"
|
||||
if ($line -match '^\[model_providers\.xixu\][ \t]*$') {
|
||||
$inXixuBlock = $true
|
||||
continue
|
||||
}
|
||||
|
||||
$filteredLines.Add($line)
|
||||
}
|
||||
|
||||
$insertIndex = 0
|
||||
while (
|
||||
$insertIndex -lt $filteredLines.Count -and (
|
||||
$filteredLines[$insertIndex].Trim().Length -eq 0 -or
|
||||
$filteredLines[$insertIndex].TrimStart().StartsWith("#")
|
||||
)
|
||||
) {
|
||||
$insertIndex++
|
||||
}
|
||||
|
||||
$outputLines = [System.Collections.Generic.List[string]]::new()
|
||||
for ($i = 0; $i -lt $insertIndex; $i++) {
|
||||
$outputLines.Add($filteredLines[$i])
|
||||
}
|
||||
|
||||
if ($outputLines.Count -gt 0 -and $outputLines[$outputLines.Count - 1].Trim().Length -ne 0) {
|
||||
$outputLines.Add("")
|
||||
}
|
||||
|
||||
$outputLines.Add('model_provider = "xixu"')
|
||||
|
||||
if ($insertIndex -lt $filteredLines.Count -and $filteredLines[$insertIndex].Trim().Length -ne 0) {
|
||||
$outputLines.Add("")
|
||||
}
|
||||
|
||||
for ($i = $insertIndex; $i -lt $filteredLines.Count; $i++) {
|
||||
$outputLines.Add($filteredLines[$i])
|
||||
}
|
||||
|
||||
while ($outputLines.Count -gt 0 -and $outputLines[$outputLines.Count - 1].Trim().Length -eq 0) {
|
||||
$outputLines.RemoveAt($outputLines.Count - 1)
|
||||
}
|
||||
|
||||
if ($outputLines.Count -gt 0) {
|
||||
$outputLines.Add("")
|
||||
}
|
||||
|
||||
$outputLines.Add('[model_providers.xixu]')
|
||||
$outputLines.Add('name = "Xi Xu''s AI Inference"')
|
||||
$outputLines.Add('base_url = "https://api.xi-xu.me/v1"')
|
||||
$outputLines.Add('env_key = "XIXU_API_KEY"')
|
||||
|
||||
$content = [string]::Join("`r`n", $outputLines) + "`r`n"
|
||||
|
||||
Set-Content -Path $ConfigFile -Value $content -Encoding UTF8
|
||||
|
||||
|
||||
@@ -24,14 +24,17 @@ TMP_FILE="$(mktemp)"
|
||||
awk '
|
||||
BEGIN {
|
||||
in_xixu = 0
|
||||
model_provider_written = 0
|
||||
}
|
||||
{
|
||||
if ($0 ~ /^[[:space:]]*model_provider[[:space:]]*=/) {
|
||||
if (!model_provider_written) {
|
||||
print "model_provider = \"xixu\""
|
||||
model_provider_written = 1
|
||||
if (in_xixu) {
|
||||
if ($0 ~ /^\[.*\][[:space:]]*$/) {
|
||||
in_xixu = 0
|
||||
} else {
|
||||
next
|
||||
}
|
||||
}
|
||||
|
||||
if ($0 ~ /^[[:space:]]*model_provider[[:space:]]*=/) {
|
||||
next
|
||||
}
|
||||
|
||||
@@ -40,18 +43,29 @@ BEGIN {
|
||||
next
|
||||
}
|
||||
|
||||
if (in_xixu && $0 ~ /^\[.*\][[:space:]]*$/) {
|
||||
in_xixu = 0
|
||||
}
|
||||
|
||||
if (!in_xixu) {
|
||||
print $0
|
||||
if (body_count == 0 && $0 ~ /^[[:space:]]*($|#)/) {
|
||||
lead[++lead_count] = $0
|
||||
} else {
|
||||
body[++body_count] = $0
|
||||
}
|
||||
}
|
||||
END {
|
||||
if (!model_provider_written) {
|
||||
for (i = 1; i <= lead_count; i++) {
|
||||
print lead[i]
|
||||
}
|
||||
|
||||
if (lead_count > 0 && lead[lead_count] !~ /^[[:space:]]*$/) {
|
||||
print ""
|
||||
print "model_provider = \"xixu\""
|
||||
}
|
||||
|
||||
print "model_provider = \"xixu\""
|
||||
|
||||
if (body_count > 0 && body[1] !~ /^[[:space:]]*$/) {
|
||||
print ""
|
||||
}
|
||||
|
||||
for (i = 1; i <= body_count; i++) {
|
||||
print body[i]
|
||||
}
|
||||
|
||||
print ""
|
||||
|
||||
@@ -92,6 +92,60 @@ class ConfigDirTests(unittest.TestCase):
|
||||
self.assertEqual(data["env"]["EXISTING"], "1")
|
||||
self.assertEqual(data["env"]["ANTHROPIC_AUTH_TOKEN"], "test-key")
|
||||
|
||||
def test_codex_ps1_writes_model_provider_at_top_level(self) -> None:
|
||||
temp_dir = tempfile.mkdtemp(prefix="provider-setup-codex-")
|
||||
self.addCleanup(lambda: shutil.rmtree(temp_dir, ignore_errors=True))
|
||||
|
||||
config_file = pathlib.Path(temp_dir) / "config.toml"
|
||||
config_file.write_text(
|
||||
"\n".join(
|
||||
[
|
||||
"# existing comment",
|
||||
"",
|
||||
"[projects.'E:\\\\github\\\\demo']",
|
||||
'trust_level = "trusted"',
|
||||
"",
|
||||
]
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
env = os.environ.copy()
|
||||
env["CODEX_HOME"] = temp_dir
|
||||
|
||||
completed = subprocess.run(
|
||||
[
|
||||
"powershell",
|
||||
"-ExecutionPolicy",
|
||||
"Bypass",
|
||||
"-File",
|
||||
str(REPO_ROOT / "codex.ps1"),
|
||||
"-ApiKey",
|
||||
"test-key",
|
||||
],
|
||||
cwd=REPO_ROOT,
|
||||
env=env,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=False,
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
completed.returncode,
|
||||
0,
|
||||
msg=f"stdout:\n{completed.stdout}\n\nstderr:\n{completed.stderr}",
|
||||
)
|
||||
|
||||
content = config_file.read_text(encoding="utf-8-sig")
|
||||
lines = content.splitlines()
|
||||
non_empty_lines = [line for line in lines if line.strip()]
|
||||
|
||||
self.assertGreaterEqual(len(non_empty_lines), 4, content)
|
||||
self.assertEqual(non_empty_lines[0], "# existing comment")
|
||||
self.assertEqual(non_empty_lines[1], 'model_provider = "xixu"')
|
||||
self.assertEqual(non_empty_lines[2], "[projects.'E:\\\\github\\\\demo']")
|
||||
self.assertIn("[model_providers.xixu]", non_empty_lines)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in new issue
Block a user