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