package service import ( "bytes" "context" "encoding/json" "fmt" "io" "net/http" "reflect" "regexp" "strings" "time" "model-gateway/model/domain" "model-gateway/model/dto" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/util/gconv" ) var SchemaMapping = &schemaMappingService{} type schemaMappingService struct{} // buildFieldDescriptions 从结构体中反射读取字段定义,构建提示词中的目标字段说明 func buildFieldDescriptions(t reflect.Type) string { var b strings.Builder for i := 0; i < t.NumField(); i++ { f := t.Field(i) jsonName := f.Tag.Get("json") desc := f.Tag.Get("dc") typeName := f.Type.String() if jsonName == "" || jsonName == "-" { continue } if b.Len() > 0 { b.WriteByte('\n') } b.WriteString("- **" + jsonName + "** (" + typeName + "): " + desc) } return b.String() } // getDomainTypeByModelType 根据模型类型返回对应的业务字段结构体反射类型 // 如果找不到匹配,返回 nil func getDomainTypeByModelType(modelType int) reflect.Type { switch modelType { case 100, 101, 102, 103, 500, 501, 502, 503: return reflect.TypeOf((*domain.ChatFieldsReq)(nil)).Elem() case 600, 601, 602, 603, 604: return reflect.TypeOf((*domain.VideoFields)(nil)).Elem() default: return nil } } // BuildSchemaMapping 根据模型类型和 schema JSON,自动构建 schema_mapping(补充已有 mapping 的缺失字段) func (s *schemaMappingService) BuildSchemaMapping(ctx context.Context, req *dto.BuildSchemaMappingReq) (res *dto.BuildSchemaMappingRes, err error) { if g.IsEmpty(req.Schema) { return nil, fmt.Errorf("schema 不能为空") } // 1. 根据模型类型获取对应的业务字段结构体 domainType := getDomainTypeByModelType(req.ModelType) if domainType == nil { return nil, fmt.Errorf("不支持的模型类型: %d", req.ModelType) } // 3. 构建 LLM 提示词 输出的 JSON 对象键是 json 字段名。每个字段的值是定位到该位置的完整点号路径。 fieldDescs := buildFieldDescriptions(domainType) systemPrompt := fmt.Sprintf(`你是一个 JSON Schema 分析助手。我提供了一个 AI API 的完整 Schema JSON 和待填充的目标结构体。 请你仔细阅读 Schema 中所有字段的名称、类型、description 描述、枚举值、约束范围等完整信息, 结合对 API 功能的理解,将目标结构体的每个字段映射到 Schema 中恰当的位置。 ## 输出格式 输出的 JSON 对象键是 json 字段名。每个字段的值有两类: 第一类(Schema 路径):若该概念在 Schema 中有直接定义位置(约束值或字段定义),输出定位到该位置的完整点号路径。若定位的是对象数组的特定元素及其属性,在路径后追加 ?实际筛选字段名=筛选值&实际值字段名=# 格式,其中 =# 标记的目标值字段名替换为 schema 中的实际字段名。 第二类(推导字符串):若该概念在 Schema 中没有直接对应的定义位置,输出根据 Schema 信息推导出的内容字符串。 ## 重要规则 1. 输出的每个字段都必须出现在 JSON 中,一个都不能少 2. 若无法从 Schema 推理出某个字段的值,就输出空字符串 "" ## 目标字段说明 %s`, fieldDescs) userPrompt := fmt.Sprintf("请分析以下 Schema JSON,生成对应的 schema_mapping:\n\n%s", req.Schema) // 4. 调用 LLM llmResp, err := callLLM(ctx, systemPrompt, userPrompt) if err != nil { return nil, err } // 5. 归一化所有路径为固定点号语法(无论模型输出哪种写法) rawMap := gconv.Map(llmResp) for k, v := range rawMap { if s, ok := v.(string); ok { rawMap[k] = normalizeSchemaPath(s) } } return &dto.BuildSchemaMappingRes{ SchemaMapping: rawMap, }, nil } // regNumIndexStar 匹配数字下标 [0]、[1] 等 var regNumIndexStar = regexp.MustCompile(`\[\d+]`) // normalizeSchemaPath 将 LLM 生成的 Schema 路径统一为固定点号语法: // - 移除模板包装字段 attrs / properties / items / defaultValue / required // - enumValues、items 及 attrs[数字] 标记上一字段为数组,补 [*] // - [数字] 下标统一转为 [*] // // 示例: // // messages.attrs.enumValues.attrs.content.enumValues?type=image_url&image_url.url=# // → messages[*].content[*]?type=image_url&image_url.url=# // choices.attrs[0].attrs.message.attrs.content → choices[*].message.content func normalizeSchemaPath(p string) string { path, suffix := p, "" if i := strings.Index(p, "?"); i >= 0 { path, suffix = p[:i], p[i:] } segs := strings.Split(path, ".") var out []string for _, seg := range segs { seg = strings.TrimSpace(seg) switch { case seg == "": continue case seg == "attrs" || seg == "properties" || seg == "defaultValue" || seg == "required": continue case seg == "enumValues" || seg == "items" || (strings.HasPrefix(seg, "attrs[") && regNumIndexStar.MatchString(seg)): markPrevAsArray(&out) continue } seg = regNumIndexStar.ReplaceAllString(seg, "[*]") out = append(out, seg) } return strings.Join(out, ".") + suffix } // markPrevAsArray 将输出序列最后一个字段标记为数组(补 [*]) func markPrevAsArray(out *[]string) { if len(*out) == 0 { return } last := (*out)[len(*out)-1] if !strings.HasSuffix(last, "[]") && !strings.HasSuffix(last, "[*]") { (*out)[len(*out)-1] = last + "[*]" } } // callLLM 调用大模型聊天接口(OpenAI 兼容格式) func callLLM(ctx context.Context, systemPrompt, userPrompt string) (string, error) { modelName := "doubao-seed-2-0-lite-260428" baseURL := "https://ark.cn-beijing.volces.com/api/v3/chat/completions" apiKey := "ark-9df744e8-a0de-4c54-9db3-18379bccd523-e6733" body := map[string]any{ "model": modelName, "messages": []map[string]string{ {"role": "system", "content": systemPrompt}, {"role": "user", "content": userPrompt}, }, "max_tokens": 2048, "temperature": 0.1, } jsonBody, err := json.Marshal(body) if err != nil { return "", fmt.Errorf("marshal request body failed: %w", err) } url := strings.TrimRight(baseURL, "/") httpReq, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonBody)) if err != nil { return "", fmt.Errorf("create request failed: %w", err) } httpReq.Header.Set("Authorization", "Bearer "+apiKey) httpReq.Header.Set("Content-Type", "application/json") client := &http.Client{Timeout: 120 * time.Second} resp, err := client.Do(httpReq) if err != nil { return "", fmt.Errorf("request failed: %w", err) } defer resp.Body.Close() respBody, err := io.ReadAll(resp.Body) if err != nil { return "", fmt.Errorf("read response failed (status=%d): %w", resp.StatusCode, err) } if resp.StatusCode != 200 { return "", fmt.Errorf("API error status=%d body=%s", resp.StatusCode, string(respBody)) } var apiResp struct { Choices []struct { Message struct { Content string `json:"content"` } `json:"message"` } `json:"choices"` Error *struct { Message string `json:"message"` } `json:"error,omitempty"` } if err = json.Unmarshal(respBody, &apiResp); err != nil { return "", fmt.Errorf("parse response failed: %s", string(respBody)) } if apiResp.Error != nil { return "", fmt.Errorf("API error: %s", apiResp.Error.Message) } if len(apiResp.Choices) == 0 { return "", fmt.Errorf("empty response") } return apiResp.Choices[0].Message.Content, nil } // extractJSONObject 从字符串中提取第一个完整的 JSON 对象({...}) func extractJSONObject(s string) string { start := strings.Index(s, "{") if start < 0 { return s } for start > 0 { ch := s[start-1] if ch == ' ' || ch == '\t' || ch == '\n' || ch == '\r' { start-- } else { break } } end := strings.LastIndex(s, "}") if end <= start { return s } snippet := s[start : end+1] snippet = strings.TrimPrefix(snippet, "```json") snippet = strings.TrimPrefix(snippet, "```") snippet = strings.TrimSuffix(snippet, "```") snippet = strings.TrimSpace(snippet) return snippet }