新增 TakeBusinessFields、WriteBusinessFields、SetByPath 与 GetByPath 等工具,支持按映射路径写入请求体与解析响应,并更新相关依赖。
260 lines
7.9 KiB
Go
260 lines
7.9 KiB
Go
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
|
||
}
|