211 lines
7.0 KiB
Go
211 lines
7.0 KiB
Go
package oci
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"time"
|
|
|
|
"github.com/oracle/oci-go-sdk/v65/generativeai"
|
|
"github.com/oracle/oci-go-sdk/v65/generativeaiinference"
|
|
|
|
"oci-portal/internal/aiwire"
|
|
)
|
|
|
|
// GenAiModel 是区域可用基础模型的摘要(管理面 ListModels)。
|
|
type GenAiModel struct {
|
|
Ocid string `json:"ocid"`
|
|
Name string `json:"name"`
|
|
Vendor string `json:"vendor"`
|
|
Caps []string `json:"capabilities"`
|
|
// Capability 是网关侧归一能力:CHAT / EMBEDDING(兼具时算 CHAT)
|
|
Capability string `json:"capability"`
|
|
// Deprecated 是 OCI 宣布的弃用时间(TimeDeprecated),nil 表示未宣布;
|
|
// 弃用后模型仍可调用,直到 Retired(TimeOnDemandRetired 按需推理退役)。
|
|
Deprecated *time.Time `json:"deprecated"`
|
|
// Retired 是按需推理退役时间;已过表示调用必 404,同步层直接剔除
|
|
Retired *time.Time `json:"retired"`
|
|
}
|
|
|
|
func (c *RealClient) genAiClient(cred Credentials, region string) (generativeai.GenerativeAiClient, error) {
|
|
gc, err := generativeai.NewGenerativeAiClientWithConfigurationProvider(provider(cred))
|
|
if err != nil {
|
|
return gc, fmt.Errorf("new generative ai client: %w", err)
|
|
}
|
|
applyProxy(&gc.BaseClient, cred)
|
|
if region != "" {
|
|
gc.SetRegion(normalizeRegion(region))
|
|
}
|
|
return gc, nil
|
|
}
|
|
|
|
func (c *RealClient) genAiInferenceClient(cred Credentials, region string) (generativeaiinference.GenerativeAiInferenceClient, error) {
|
|
ic, err := generativeaiinference.NewGenerativeAiInferenceClientWithConfigurationProvider(provider(cred))
|
|
if err != nil {
|
|
return ic, fmt.Errorf("new generative ai inference client: %w", err)
|
|
}
|
|
applyProxy(&ic.BaseClient, cred)
|
|
if region != "" {
|
|
ic.SetRegion(normalizeRegion(region))
|
|
}
|
|
return ic, nil
|
|
}
|
|
|
|
// ListGenAiModels 实现 Client:列出区域的 CHAT / EMBEDDING 能力 ACTIVE 基础模型(去重按名称)。
|
|
func (c *RealClient) ListGenAiModels(ctx context.Context, cred Credentials, region string) ([]GenAiModel, error) {
|
|
gc, err := c.genAiClient(cred, region)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
resp, err := gc.ListModels(ctx, generativeai.ListModelsRequest{CompartmentId: &cred.TenancyOCID})
|
|
if err != nil {
|
|
return nil, fmt.Errorf("list genai models: %w", err)
|
|
}
|
|
return dedupGenAiModels(resp.Items, time.Now()), nil
|
|
}
|
|
|
|
// dedupGenAiModels 压平并按名称去重;同名多条目时优先保留不含 FINE_TUNE 能力的条目
|
|
// (微调基座条目在部分区域不支持按需调用,缓存其 OCID 会导致调用 400)。
|
|
func dedupGenAiModels(items []generativeai.ModelSummary, now time.Time) []GenAiModel {
|
|
seen := map[string]int{}
|
|
var out []GenAiModel
|
|
for _, m := range items {
|
|
gm, ok := toGenAiModel(m, now)
|
|
if !ok {
|
|
continue
|
|
}
|
|
if i, dup := seen[gm.Name]; dup {
|
|
if hasFineTune(out[i].Caps) && !hasFineTune(gm.Caps) {
|
|
out[i] = gm
|
|
}
|
|
continue
|
|
}
|
|
seen[gm.Name] = len(out)
|
|
out = append(out, gm)
|
|
}
|
|
return out
|
|
}
|
|
|
|
// hasFineTune 判断能力列表是否含 FINE_TUNE(微调基座条目)。
|
|
func hasFineTune(caps []string) bool {
|
|
for _, c := range caps {
|
|
if c == string(generativeai.ModelCapabilityFineTune) {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
// toGenAiModel 压平模型摘要;无归一能力、无名称或按需推理已退役
|
|
// (调用必 404,ListModels 仍会返回且 state 为 ACTIVE)时返回 false 不入池。
|
|
func toGenAiModel(m generativeai.ModelSummary, now time.Time) (GenAiModel, bool) {
|
|
capability := modelCapability(m)
|
|
name := deref(m.DisplayName)
|
|
if capability == "" || name == "" {
|
|
return GenAiModel{}, false
|
|
}
|
|
if m.TimeOnDemandRetired != nil && m.TimeOnDemandRetired.Time.Before(now) {
|
|
return GenAiModel{}, false
|
|
}
|
|
gm := GenAiModel{Ocid: deref(m.Id), Name: name, Vendor: deref(m.Vendor),
|
|
Caps: capStrings(m.Capabilities), Capability: capability}
|
|
if m.TimeDeprecated != nil {
|
|
gm.Deprecated = &m.TimeDeprecated.Time
|
|
}
|
|
if m.TimeOnDemandRetired != nil {
|
|
gm.Retired = &m.TimeOnDemandRetired.Time
|
|
}
|
|
return gm, true
|
|
}
|
|
|
|
// modelCapability 归一模型能力:ACTIVE 且具 CHAT / TEXT_EMBEDDINGS 才纳入(兼具时算 CHAT)。
|
|
func modelCapability(m generativeai.ModelSummary) string {
|
|
if m.LifecycleState != generativeai.ModelLifecycleStateActive {
|
|
return ""
|
|
}
|
|
capability := ""
|
|
for _, cap := range m.Capabilities {
|
|
switch cap {
|
|
case generativeai.ModelCapabilityChat:
|
|
return "CHAT"
|
|
case generativeai.ModelCapabilityTextEmbeddings:
|
|
capability = "EMBEDDING"
|
|
case generativeai.ModelCapabilityTextRerank:
|
|
capability = "RERANK"
|
|
case generativeai.ModelCapabilityEnum("TEXT_TO_AUDIO"):
|
|
// SDK v65.120 尚无该枚举常量,按原始字符串匹配(xai.grok-tts)
|
|
capability = "TTS"
|
|
}
|
|
}
|
|
return capability
|
|
}
|
|
|
|
func capStrings(caps []generativeai.ModelCapabilityEnum) []string {
|
|
out := make([]string, 0, len(caps))
|
|
for _, c := range caps {
|
|
out = append(out, string(c))
|
|
}
|
|
return out
|
|
}
|
|
|
|
// GenAiProbeChat 实现 Client:经 OpenAI 兼容面(直通同链路)发一次极小请求探测渠道;
|
|
// 配额探测专用的最小聊天,返回 HTTP 状态码。max_output_tokens 取 16:
|
|
// openai.gpt-oss 系列要求 >=16,其余模型均兼容,成本差异可忽略。
|
|
func (c *RealClient) GenAiProbeChat(ctx context.Context, cred Credentials, region, modelOcid, modelName string) (int, error) {
|
|
body, err := json.Marshal(map[string]any{"model": modelName, "input": "hi",
|
|
"max_output_tokens": 16, "store": false})
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
// 探测追求快速失败,沿用 SDK 默认量级的 60s 预算即可
|
|
if _, err = c.GenAiCompatResponses(ctx, cred, region, body, 60*time.Second); err == nil {
|
|
return http.StatusOK, nil
|
|
}
|
|
if status, ok := ServiceStatus(err); ok {
|
|
return status, err
|
|
}
|
|
return 0, err
|
|
}
|
|
|
|
// sdkUsageToIR 把 SDK 用量转为内部记账结构(缓存命中挂 details,仅命中时出现)。
|
|
func sdkUsageToIR(u *generativeaiinference.Usage) *aiwire.Usage {
|
|
if u == nil {
|
|
return nil
|
|
}
|
|
out := &aiwire.Usage{}
|
|
if u.PromptTokens != nil {
|
|
out.PromptTokens = *u.PromptTokens
|
|
}
|
|
if u.CompletionTokens != nil {
|
|
out.CompletionTokens = *u.CompletionTokens
|
|
}
|
|
if u.TotalTokens != nil {
|
|
out.TotalTokens = *u.TotalTokens
|
|
}
|
|
if u.PromptTokensDetails != nil && u.PromptTokensDetails.CachedTokens != nil && *u.PromptTokensDetails.CachedTokens > 0 {
|
|
out.PromptTokensDetails = &aiwire.PromptTokensDetails{CachedTokens: *u.PromptTokensDetails.CachedTokens}
|
|
}
|
|
return out
|
|
}
|
|
|
|
// GenAiEmbed 实现 Client:文本向量化(on-demand serving);dimensions 透传 OutputDimensions。
|
|
func (c *RealClient) GenAiEmbed(ctx context.Context, cred Credentials, region, modelOcid string, inputs []string, dimensions *int) ([][]float32, *aiwire.Usage, error) {
|
|
ic, err := c.genAiInferenceClient(cred, region)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
resp, err := ic.EmbedText(ctx, generativeaiinference.EmbedTextRequest{
|
|
EmbedTextDetails: generativeaiinference.EmbedTextDetails{
|
|
CompartmentId: &cred.TenancyOCID,
|
|
ServingMode: generativeaiinference.OnDemandServingMode{ModelId: &modelOcid},
|
|
Inputs: inputs,
|
|
OutputDimensions: dimensions,
|
|
},
|
|
})
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("genai embed: %w", err)
|
|
}
|
|
return resp.Embeddings, sdkUsageToIR(resp.Usage), nil
|
|
}
|