返回 DeepSeek-Reasonix
provider_endpoint_rewrite.go
根目录 / internal / config / provider_endpoint_rewrite.go
1 package config
2
3 import (
4 "fmt"
5 "os"
6 "slices"
7 "sort"
8 "strconv"
9 "strings"
10
11 "github.com/BurntSushi/toml"
12
13 "reasonix/internal/fileutil"
14 fileencoding "reasonix/internal/fileutil/encoding"
15 )
16
17 // repairProviderEndpointContractsFileLocked performs a narrow lexical edit for
18 // a caller that already owns LockConfigFileEdits. It preserves comments,
19 // unknown provider fields, inline tables, encoding and file permissions.
20 func repairProviderEndpointContractsFileLocked(path string) ([]ProviderEndpointRepair, error) {
21 resolved, exists, err := statConfigPath(path)
22 if err != nil || !exists {
23 return nil, err
24 }
25 info, err := os.Stat(resolved)
26 if err != nil {
27 return nil, err
28 }
29 rawBytes, err := os.ReadFile(resolved)
30 if err != nil {
31 return nil, err
32 }
33 encoding, detected := fileencoding.Detect(rawBytes)
34 raw := fileencoding.Decode(detected, encoding)
35 next, repairs, err := rewriteProviderEndpointContracts(string(raw))
36 if err != nil || len(repairs) == 0 {
37 return repairs, err
38 }
39 encoded, err := fileencoding.Encode(next, encoding)
40 if err != nil {
41 return nil, err
42 }
43 if err := fileutil.AtomicWriteFile(resolved, encoded, info.Mode().Perm()); err != nil {
44 return nil, err
45 }
46 return repairs, nil
47 }
48
49 func rewriteProviderEndpointContracts(raw string) (string, []ProviderEndpointRepair, error) {
50 var decoded struct {
51 Providers []ProviderEntry `toml:"providers"`
52 }
53 if _, err := toml.Decode(raw, &decoded); err != nil {
54 return raw, nil, err
55 }
56 repairs := make([]ProviderEndpointRepair, 0)
57 repaired := make([]bool, len(decoded.Providers))
58 for i := range decoded.Providers {
59 if repair, changed := RepairProviderEndpointContract(&decoded.Providers[i]); changed {
60 repairs = append(repairs, *repair)
61 repaired[i] = true
62 }
63 }
64 if len(repairs) == 0 {
65 return raw, nil, nil
66 }
67
68 lines := strings.Split(raw, "\n")
69 blocks := providerTOMLBlocks(lines)
70 if len(blocks) == len(decoded.Providers) {
71 var err error
72 for i := range slices.Backward(decoded.Providers) {
73 if !repaired[i] {
74 continue
75 }
76 lines, err = rewriteProviderEndpointBlock(lines, blocks[i], decoded.Providers[i])
77 if err != nil {
78 return raw, nil, err
79 }
80 }
81 return strings.Join(lines, "\n"), repairs, nil
82 }
83
84 inlineBlocks, err := providerTOMLInlineBlocks(raw)
85 if err != nil || len(inlineBlocks) != len(decoded.Providers) {
86 return raw, nil, fmt.Errorf("repair provider endpoint: could not map provider tables safely")
87 }
88 replacements := make([]tomlReplacement, 0, len(repairs)*7)
89 for i := range decoded.Providers {
90 if !repaired[i] {
91 continue
92 }
93 blockReplacements, err := providerEndpointInlineReplacements(raw, inlineBlocks[i], decoded.Providers[i])
94 if err != nil {
95 return raw, nil, err
96 }
97 replacements = append(replacements, blockReplacements...)
98 }
99 return applyTOMLReplacements(raw, replacements), repairs, nil
100 }
101
102 func rewriteProviderEndpointBlock(lines []string, block providerTOMLBlock, entry ProviderEntry) ([]string, error) {
103 foundKind, foundBase := false, false
104 foundAuth, foundMode := false, false
105 state := tomlOutside
106 for i := block.start + 1; i < block.end; i++ {
107 if state != tomlOutside {
108 state = advanceTOMLStringState(state, lines[i])
109 continue
110 }
111 nextState := advanceTOMLStringState(tomlOutside, lines[i])
112 if nextState != tomlOutside {
113 state = nextState
114 continue
115 }
116 key, _, ok := tomlKeyValue(lines[i])
117 if !ok {
118 state = nextState
119 continue
120 }
121 switch strings.Trim(key, `"'`) {
122 case "kind":
123 lines[i] = replaceTOMLStringAssignment(lines[i], entry.Kind)
124 foundKind = true
125 case "base_url":
126 lines[i] = replaceTOMLStringAssignment(lines[i], entry.BaseURL)
127 foundBase = true
128 case "request_url", "chat_url":
129 lines[i] = replaceTOMLStringAssignment(lines[i], "")
130 case "auth_header":
131 lines[i] = replaceTOMLScalarAssignment(lines[i], strconv.FormatBool(entry.AuthHeader))
132 foundAuth = true
133 case "responses_mode":
134 foundMode = true
135 if entry.Kind == "responses" && entry.ResponsesMode != "" {
136 lines[i] = replaceTOMLStringAssignment(lines[i], entry.ResponsesMode)
137 } else {
138 lines[i] = preservedTOMLLineComment(lines[i])
139 }
140 case "responses_stateful":
141 lines[i] = preservedTOMLLineComment(lines[i])
142 }
143 state = nextState
144 }
145 insert := make([]string, 0, 4)
146 if !foundKind {
147 insert = append(insert, "kind = "+strconv.Quote(entry.Kind))
148 }
149 if !foundBase {
150 insert = append(insert, "base_url = "+strconv.Quote(entry.BaseURL))
151 }
152 if entry.AuthHeader && !foundAuth {
153 insert = append(insert, "auth_header = true")
154 }
155 if entry.Kind == "responses" && entry.ResponsesMode != "" && !foundMode {
156 insert = append(insert, "responses_mode = "+strconv.Quote(entry.ResponsesMode))
157 }
158 if len(insert) == 0 {
159 return lines, nil
160 }
161 if block.end > 0 && strings.HasSuffix(lines[block.end-1], "\r") {
162 for i := range insert {
163 insert[i] += "\r"
164 }
165 }
166 lines = append(lines, make([]string, len(insert))...)
167 copy(lines[block.end+len(insert):], lines[block.end:len(lines)-len(insert)])
168 copy(lines[block.end:], insert)
169 return lines, nil
170 }
171
172 func preservedTOMLLineComment(line string) string {
173 carriageReturn := strings.HasSuffix(line, "\r")
174 line = strings.TrimSuffix(line, "\r")
175 comment := tomlInlineCommentIndex(line)
176 if comment < 0 {
177 if carriageReturn {
178 return "\r"
179 }
180 return ""
181 }
182 indentLen := len(line) - len(strings.TrimLeft(line, " \t"))
183 next := line[:indentLen] + strings.TrimLeft(line[comment:], " \t")
184 if carriageReturn {
185 next += "\r"
186 }
187 return next
188 }
189
190 func providerEndpointInlineReplacements(raw string, block providerTOMLInlineBlock, entry ProviderEntry) ([]tomlReplacement, error) {
191 removeKeys := map[string]bool{"responses_stateful": true}
192 if entry.Kind != "responses" || entry.ResponsesMode == "" {
193 removeKeys["responses_mode"] = true
194 }
195 replacements := make([]tomlReplacement, 0, 8)
196 if kind, ok := block.fields["kind"]; ok {
197 replacements = append(replacements, tomlReplacement{start: kind.valueStart, end: kind.valueEnd, value: strconv.Quote(entry.Kind)})
198 }
199 if base, ok := block.fields["base_url"]; ok {
200 replacements = append(replacements, tomlReplacement{start: base.valueStart, end: base.valueEnd, value: strconv.Quote(entry.BaseURL)})
201 }
202 for _, key := range []string{"request_url", "chat_url"} {
203 if field, ok := block.fields[key]; ok {
204 replacements = append(replacements, tomlReplacement{start: field.valueStart, end: field.valueEnd, value: strconv.Quote("")})
205 }
206 }
207 if field, ok := block.fields["auth_header"]; ok {
208 replacements = append(replacements, tomlReplacement{start: field.valueStart, end: field.valueEnd, value: strconv.FormatBool(entry.AuthHeader)})
209 }
210 if field, ok := block.fields["responses_mode"]; ok && !removeKeys["responses_mode"] {
211 replacements = append(replacements, tomlReplacement{start: field.valueStart, end: field.valueEnd, value: strconv.Quote(entry.ResponsesMode)})
212 }
213 replacements = append(replacements, inlineProviderFieldRemovals(block, removeKeys)...)
214
215 var additions []string
216 if _, ok := block.fields["kind"]; !ok {
217 additions = append(additions, "kind = "+strconv.Quote(entry.Kind))
218 }
219 if _, ok := block.fields["base_url"]; !ok {
220 additions = append(additions, "base_url = "+strconv.Quote(entry.BaseURL))
221 }
222 if entry.AuthHeader {
223 if _, ok := block.fields["auth_header"]; !ok {
224 additions = append(additions, "auth_header = true")
225 }
226 }
227 if entry.Kind == "responses" && entry.ResponsesMode != "" {
228 if _, ok := block.fields["responses_mode"]; !ok {
229 additions = append(additions, "responses_mode = "+strconv.Quote(entry.ResponsesMode))
230 }
231 }
232 if len(additions) > 0 {
233 replacements = append(replacements, tomlReplacement{
234 start: block.end,
235 end: block.end,
236 value: ", " + strings.Join(additions, ", "),
237 })
238 }
239 return replacements, nil
240 }
241
242 func inlineProviderFieldRemovals(block providerTOMLInlineBlock, removeKeys map[string]bool) []tomlReplacement {
243 indexSet := make(map[int]bool)
244 for key := range removeKeys {
245 if field, ok := block.fields[key]; ok {
246 indexSet[field.segment] = true
247 }
248 }
249 if len(indexSet) == 0 {
250 return nil
251 }
252 indexes := make([]int, 0, len(indexSet))
253 for index := range indexSet {
254 indexes = append(indexes, index)
255 }
256 sort.Ints(indexes)
257 var replacements []tomlReplacement
258 for pos := 0; pos < len(indexes); {
259 first, last := indexes[pos], indexes[pos]
260 for pos+1 < len(indexes) && indexes[pos+1] == last+1 {
261 pos++
262 last = indexes[pos]
263 }
264 var start, end int
265 if last < len(block.segments)-1 {
266 start = block.segments[first][0]
267 end = block.segments[last+1][0]
268 } else if first > 0 {
269 start = block.segments[first-1][1]
270 end = block.segments[last][1]
271 }
272 if end > start {
273 replacements = append(replacements, tomlReplacement{start: start, end: end})
274 }
275 pos++
276 }
277 return replacements
278 }
279
279 lines GO