Files
model-gateway/service/schema_mapping_service.go

261 lines
8.1 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 兼容格式)。
// 模型地址/密钥走配置 schemaMapping 段,本地开发无配置时用默认值兜底。
func callLLM(ctx context.Context, systemPrompt, userPrompt string) (string, error) {
modelName := g.Cfg().MustGet(ctx, "schemaMapping.modelName", "doubao-seed-2-0-lite-260428").String()
baseURL := g.Cfg().MustGet(ctx, "schemaMapping.baseUrl", "https://ark.cn-beijing.volces.com/api/v3/chat/completions").String()
apiKey := g.Cfg().MustGet(ctx, "schemaMapping.apiKey", "ark-9df744e8-a0de-4c54-9db3-18379bccd523-e6733").String()
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
}