openai/openai-go

Public

mirrored from https://github.com/openai/openai-goAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
release-please--branches--main--changes--next

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

internal/apiform/encoder.go

493lines · modecode

1package apiform
2
3import (
4 "fmt"
5 "io"
6 "mime/multipart"
7 "net/textproto"
8 "path"
9 "reflect"
10 "sort"
11 "strconv"
12 "strings"
13 "sync"
14 "time"
15
16 "github.com/openai/openai-go/v3/packages/param"
17)
18
19var encoders sync.Map // map[encoderEntry]encoderFunc
20
21func Marshal(value any, writer *multipart.Writer) error {
22 e := &encoder{
23 dateFormat: time.RFC3339,
24 arrayFmt: "brackets",
25 }
26 return e.marshal(value, writer)
27}
28
29func MarshalRoot(value any, writer *multipart.Writer) error {
30 e := &encoder{
31 root: true,
32 dateFormat: time.RFC3339,
33 arrayFmt: "brackets",
34 }
35 return e.marshal(value, writer)
36}
37
38func MarshalWithSettings(value any, writer *multipart.Writer, arrayFormat string) error {
39 e := &encoder{
40 arrayFmt: arrayFormat,
41 dateFormat: time.RFC3339,
42 }
43 return e.marshal(value, writer)
44}
45
46type encoder struct {
47 arrayFmt string
48 dateFormat string
49 root bool
50}
51
52type encoderFunc func(key string, value reflect.Value, writer *multipart.Writer) error
53
54type encoderField struct {
55 tag parsedStructTag
56 fn encoderFunc
57 idx []int
58}
59
60type encoderEntry struct {
61 typ reflect.Type
62 dateFormat string
63 arrayFmt string
64 root bool
65}
66
67func (e *encoder) marshal(value any, writer *multipart.Writer) error {
68 val := reflect.ValueOf(value)
69 if !val.IsValid() {
70 return nil
71 }
72 typ := val.Type()
73 enc := e.typeEncoder(typ)
74 return enc("", val, writer)
75}
76
77func (e *encoder) typeEncoder(t reflect.Type) encoderFunc {
78 entry := encoderEntry{
79 typ: t,
80 dateFormat: e.dateFormat,
81 arrayFmt: e.arrayFmt,
82 root: e.root,
83 }
84
85 if fi, ok := encoders.Load(entry); ok {
86 return fi.(encoderFunc)
87 }
88
89 // To deal with recursive types, populate the map with an
90 // indirect func before we build it. This type waits on the
91 // real func (f) to be ready and then calls it. This indirect
92 // func is only used for recursive types.
93 var (
94 wg sync.WaitGroup
95 f encoderFunc
96 )
97 wg.Add(1)
98 fi, loaded := encoders.LoadOrStore(entry, encoderFunc(func(key string, v reflect.Value, writer *multipart.Writer) error {
99 wg.Wait()
100 return f(key, v, writer)
101 }))
102 if loaded {
103 return fi.(encoderFunc)
104 }
105
106 // Compute the real encoder and replace the indirect func with it.
107 f = e.newTypeEncoder(t)
108 wg.Done()
109 encoders.Store(entry, f)
110 return f
111}
112
113func (e *encoder) newTypeEncoder(t reflect.Type) encoderFunc {
114 if t.ConvertibleTo(reflect.TypeOf(time.Time{})) {
115 return e.newTimeTypeEncoder()
116 }
117 if t.Implements(reflect.TypeOf((*io.Reader)(nil)).Elem()) {
118 return e.newReaderTypeEncoder()
119 }
120 e.root = false
121 switch t.Kind() {
122 case reflect.Pointer:
123 inner := t.Elem()
124
125 innerEncoder := e.typeEncoder(inner)
126 return func(key string, v reflect.Value, writer *multipart.Writer) error {
127 if !v.IsValid() || v.IsNil() {
128 return nil
129 }
130 return innerEncoder(key, v.Elem(), writer)
131 }
132 case reflect.Struct:
133 return e.newStructTypeEncoder(t)
134 case reflect.Slice, reflect.Array:
135 return e.newArrayTypeEncoder(t)
136 case reflect.Map:
137 return e.newMapEncoder(t)
138 case reflect.Interface:
139 return e.newInterfaceEncoder()
140 default:
141 return e.newPrimitiveTypeEncoder(t)
142 }
143}
144
145func (e *encoder) newPrimitiveTypeEncoder(t reflect.Type) encoderFunc {
146 switch t.Kind() {
147 // Note that we could use `gjson` to encode these types but it would complicate our
148 // code more and this current code shouldn't cause any issues
149 case reflect.String:
150 return func(key string, v reflect.Value, writer *multipart.Writer) error {
151 return writer.WriteField(key, v.String())
152 }
153 case reflect.Bool:
154 return func(key string, v reflect.Value, writer *multipart.Writer) error {
155 if v.Bool() {
156 return writer.WriteField(key, "true")
157 }
158 return writer.WriteField(key, "false")
159 }
160 case reflect.Int, reflect.Int16, reflect.Int32, reflect.Int64:
161 return func(key string, v reflect.Value, writer *multipart.Writer) error {
162 return writer.WriteField(key, strconv.FormatInt(v.Int(), 10))
163 }
164 case reflect.Uint, reflect.Uint16, reflect.Uint32, reflect.Uint64:
165 return func(key string, v reflect.Value, writer *multipart.Writer) error {
166 return writer.WriteField(key, strconv.FormatUint(v.Uint(), 10))
167 }
168 case reflect.Float32:
169 return func(key string, v reflect.Value, writer *multipart.Writer) error {
170 return writer.WriteField(key, strconv.FormatFloat(v.Float(), 'f', -1, 32))
171 }
172 case reflect.Float64:
173 return func(key string, v reflect.Value, writer *multipart.Writer) error {
174 return writer.WriteField(key, strconv.FormatFloat(v.Float(), 'f', -1, 64))
175 }
176 default:
177 return func(key string, v reflect.Value, writer *multipart.Writer) error {
178 return fmt.Errorf("unknown type received at primitive encoder: %s", t.String())
179 }
180 }
181}
182
183func (e *encoder) newArrayTypeEncoder(t reflect.Type) encoderFunc {
184 itemEncoder := e.typeEncoder(t.Elem())
185 keyFn := e.arrayKeyEncoder()
186 if e.arrayFmt == "comma" {
187 return func(key string, v reflect.Value, writer *multipart.Writer) error {
188 if v.Len() == 0 {
189 return nil
190 }
191 elements := make([]string, v.Len())
192 for i := 0; i < v.Len(); i++ {
193 elements[i] = fmt.Sprint(v.Index(i).Interface())
194 }
195 return writer.WriteField(key, strings.Join(elements, ","))
196 }
197 }
198 return func(key string, v reflect.Value, writer *multipart.Writer) error {
199 if keyFn == nil {
200 return fmt.Errorf("apiform: unsupported array format")
201 }
202 for i := 0; i < v.Len(); i++ {
203 err := itemEncoder(keyFn(key, i), v.Index(i), writer)
204 if err != nil {
205 return err
206 }
207 }
208 return nil
209 }
210}
211
212func (e *encoder) newStructTypeEncoder(t reflect.Type) encoderFunc {
213 if t.Implements(reflect.TypeOf((*param.Optional)(nil)).Elem()) {
214 return e.newRichFieldTypeEncoder(t)
215 }
216
217 for i := 0; i < t.NumField(); i++ {
218 if t.Field(i).Type == paramUnionType && t.Field(i).Anonymous {
219 return e.newStructUnionTypeEncoder(t)
220 }
221 }
222
223 encoderFields := []encoderField{}
224 extraEncoder := (*encoderField)(nil)
225
226 // This helper allows us to recursively collect field encoders into a flat
227 // array. The parameter `index` keeps track of the access patterns necessary
228 // to get to some field.
229 var collectEncoderFields func(r reflect.Type, index []int)
230 collectEncoderFields = func(r reflect.Type, index []int) {
231 for i := 0; i < r.NumField(); i++ {
232 idx := append(index, i)
233 field := t.FieldByIndex(idx)
234 if !field.IsExported() {
235 continue
236 }
237 // If this is an embedded struct, traverse one level deeper to extract
238 // the field and get their encoders as well.
239 if field.Anonymous {
240 collectEncoderFields(field.Type, idx)
241 continue
242 }
243 // If json tag is not present, then we skip, which is intentionally
244 // different behavior from the stdlib.
245 ptag, ok := parseFormStructTag(field)
246 if !ok {
247 continue
248 }
249 // We only want to support unexported field if they're tagged with
250 // `extras` because that field shouldn't be part of the public API. We
251 // also want to only keep the top level extras
252 if ptag.extras && len(index) == 0 {
253 extraEncoder = &encoderField{ptag, e.typeEncoder(field.Type.Elem()), idx}
254 continue
255 }
256 if ptag.name == "-" || ptag.name == "" {
257 continue
258 }
259
260 dateFormat, ok := parseFormatStructTag(field)
261 oldFormat := e.dateFormat
262 if ok {
263 switch dateFormat {
264 case "date-time":
265 e.dateFormat = time.RFC3339
266 case "date":
267 e.dateFormat = "2006-01-02"
268 }
269 }
270
271 var encoderFn encoderFunc
272 if ptag.omitzero {
273 typeEncoderFn := e.typeEncoder(field.Type)
274 encoderFn = func(key string, value reflect.Value, writer *multipart.Writer) error {
275 if value.IsZero() {
276 return nil
277 }
278 return typeEncoderFn(key, value, writer)
279 }
280 } else if ptag.defaultValue != nil {
281 typeEncoderFn := e.typeEncoder(field.Type)
282 encoderFn = func(key string, value reflect.Value, writer *multipart.Writer) error {
283 if value.IsZero() {
284 return typeEncoderFn(key, reflect.ValueOf(ptag.defaultValue), writer)
285 }
286 return typeEncoderFn(key, value, writer)
287 }
288 } else {
289 encoderFn = e.typeEncoder(field.Type)
290 }
291 encoderFields = append(encoderFields, encoderField{ptag, encoderFn, idx})
292 e.dateFormat = oldFormat
293 }
294 }
295 collectEncoderFields(t, []int{})
296
297 // Ensure deterministic output by sorting by lexicographic order
298 sort.Slice(encoderFields, func(i, j int) bool {
299 return encoderFields[i].tag.name < encoderFields[j].tag.name
300 })
301
302 return func(key string, value reflect.Value, writer *multipart.Writer) error {
303 keyFn := e.objKeyEncoder(key)
304 for _, ef := range encoderFields {
305 field := value.FieldByIndex(ef.idx)
306 err := ef.fn(keyFn(ef.tag.name), field, writer)
307 if err != nil {
308 return err
309 }
310 }
311
312 if extraEncoder != nil {
313 err := e.encodeMapEntries(key, value.FieldByIndex(extraEncoder.idx), writer)
314 if err != nil {
315 return err
316 }
317 }
318
319 return nil
320 }
321}
322
323var paramUnionType = reflect.TypeOf((*param.APIUnion)(nil)).Elem()
324
325func (e *encoder) newStructUnionTypeEncoder(t reflect.Type) encoderFunc {
326 var fieldEncoders []encoderFunc
327 for i := 0; i < t.NumField(); i++ {
328 field := t.Field(i)
329 if field.Type == paramUnionType && field.Anonymous {
330 fieldEncoders = append(fieldEncoders, nil)
331 continue
332 }
333 fieldEncoders = append(fieldEncoders, e.typeEncoder(field.Type))
334 }
335
336 return func(key string, value reflect.Value, writer *multipart.Writer) error {
337 for i := 0; i < t.NumField(); i++ {
338 if value.Field(i).Type() == paramUnionType {
339 continue
340 }
341 if !value.Field(i).IsZero() {
342 return fieldEncoders[i](key, value.Field(i), writer)
343 }
344 }
345 return fmt.Errorf("apiform: union %s has no field set", t.String())
346 }
347}
348
349func (e *encoder) newTimeTypeEncoder() encoderFunc {
350 format := e.dateFormat
351 return func(key string, value reflect.Value, writer *multipart.Writer) error {
352 return writer.WriteField(key, value.Convert(reflect.TypeOf(time.Time{})).Interface().(time.Time).Format(format))
353 }
354}
355
356func (e encoder) newInterfaceEncoder() encoderFunc {
357 return func(key string, value reflect.Value, writer *multipart.Writer) error {
358 value = value.Elem()
359 if !value.IsValid() {
360 return nil
361 }
362 return e.typeEncoder(value.Type())(key, value, writer)
363 }
364}
365
366var quoteEscaper = strings.NewReplacer("\\", "\\\\", `"`, "\\\"")
367
368func escapeQuotes(s string) string {
369 return quoteEscaper.Replace(s)
370}
371
372func (e *encoder) newReaderTypeEncoder() encoderFunc {
373 return func(key string, value reflect.Value, writer *multipart.Writer) error {
374 reader, ok := value.Convert(reflect.TypeOf((*io.Reader)(nil)).Elem()).Interface().(io.Reader)
375 if !ok {
376 return nil
377 }
378 filename := "anonymous_file"
379 contentType := "application/octet-stream"
380 if named, ok := reader.(interface{ Filename() string }); ok {
381 filename = named.Filename()
382 } else if named, ok := reader.(interface{ Name() string }); ok {
383 filename = path.Base(named.Name())
384 }
385 if typed, ok := reader.(interface{ ContentType() string }); ok {
386 contentType = typed.ContentType()
387 }
388
389 // Below is taken almost 1-for-1 from [multipart.CreateFormFile]
390 h := make(textproto.MIMEHeader)
391 h.Set("Content-Disposition", fmt.Sprintf(`form-data; name="%s"; filename="%s"`, escapeQuotes(key), escapeQuotes(filename)))
392 h.Set("Content-Type", contentType)
393 filewriter, err := writer.CreatePart(h)
394 if err != nil {
395 return err
396 }
397 _, err = io.Copy(filewriter, reader)
398 return err
399 }
400}
401
402func (e encoder) arrayKeyEncoder() func(string, int) string {
403 var keyFn func(string, int) string
404 switch e.arrayFmt {
405 case "comma", "repeat":
406 keyFn = func(k string, _ int) string { return k }
407 case "brackets":
408 keyFn = func(key string, _ int) string { return key + "[]" }
409 case "indices:dots":
410 keyFn = func(k string, i int) string {
411 if k == "" {
412 return strconv.Itoa(i)
413 }
414 return k + "." + strconv.Itoa(i)
415 }
416 case "indices:brackets":
417 keyFn = func(k string, i int) string {
418 if k == "" {
419 return strconv.Itoa(i)
420 }
421 return k + "[" + strconv.Itoa(i) + "]"
422 }
423 }
424 return keyFn
425}
426
427func (e encoder) objKeyEncoder(parent string) func(string) string {
428 if parent == "" {
429 return func(child string) string { return child }
430 }
431 switch e.arrayFmt {
432 case "brackets":
433 return func(child string) string { return parent + "[" + child + "]" }
434 default:
435 return func(child string) string { return parent + "." + child }
436 }
437}
438
439// Given a []byte of json (may either be an empty object or an object that already contains entries)
440// encode all of the entries in the map to the json byte array.
441func (e *encoder) encodeMapEntries(key string, v reflect.Value, writer *multipart.Writer) error {
442 type mapPair struct {
443 key string
444 value reflect.Value
445 }
446
447 pairs := []mapPair{}
448
449 iter := v.MapRange()
450 for iter.Next() {
451 if iter.Key().Type().Kind() == reflect.String {
452 pairs = append(pairs, mapPair{key: iter.Key().String(), value: iter.Value()})
453 } else {
454 return fmt.Errorf("cannot encode a map with a non string key")
455 }
456 }
457
458 // Ensure deterministic output
459 sort.Slice(pairs, func(i, j int) bool {
460 return pairs[i].key < pairs[j].key
461 })
462
463 elementEncoder := e.typeEncoder(v.Type().Elem())
464 keyFn := e.objKeyEncoder(key)
465 for _, p := range pairs {
466 err := elementEncoder(keyFn(p.key), p.value, writer)
467 if err != nil {
468 return err
469 }
470 }
471
472 return nil
473}
474
475func (e *encoder) newMapEncoder(_ reflect.Type) encoderFunc {
476 return func(key string, value reflect.Value, writer *multipart.Writer) error {
477 return e.encodeMapEntries(key, value, writer)
478 }
479}
480
481func WriteExtras(writer *multipart.Writer, extras map[string]any) (err error) {
482 for k, v := range extras {
483 str, ok := v.(string)
484 if !ok {
485 break
486 }
487 err = writer.WriteField(k, str)
488 if err != nil {
489 break
490 }
491 }
492 return err
493}
494