Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 13 additions & 3 deletions internal/dao/tenant_model.go
Original file line number Diff line number Diff line change
Expand Up @@ -66,10 +66,20 @@ func (dao *TenantModelDAO) GetModelByProviderIDAndInstanceIDAndModelName(provide
return &model, nil
}

// GetModelsByInstanceID get all models by instance ID
func (dao *TenantModelDAO) GetModelsByInstanceID(instanceID string) ([]*entity.TenantModel, error) {
// GetModelByProviderIDAndInstanceIDAndModelTypeAndModelName gets a model by provider ID, instance ID, model type and model name
func (dao *TenantModelDAO) GetModelByProviderIDAndInstanceIDAndModelTypeAndModelName(providerID, instanceID, modelType, modelName string) (*entity.TenantModel, error) {
var model entity.TenantModel
err := DB.Where("provider_id = ? AND instance_id = ? AND model_type = ? AND model_name = ?", providerID, instanceID, modelType, modelName).First(&model).Error
if err != nil {
return nil, err
}
return &model, nil
}

// GetModelsByInstanceIDs get all models by instance IDs
func (dao *TenantModelDAO) GetModelsByInstanceIDs(instanceIDs []string) ([]*entity.TenantModel, error) {
var models []*entity.TenantModel
err := DB.Where("instance_id = ?", instanceID).Find(&models).Error
err := DB.Where("instance_id IN ?", instanceIDs).Find(&models).Error
if err != nil {
return nil, err
}
Expand Down
19 changes: 10 additions & 9 deletions internal/dao/tenant_model_instance.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,15 +32,6 @@ func (dao *TenantModelInstanceDAO) Create(instance *entity.TenantModelInstance)
return DB.Create(instance).Error
}

func (dao *TenantModelInstanceDAO) GetAllInstancesByProviderID(providerID string) ([]*entity.TenantModelInstance, error) {
var instances []*entity.TenantModelInstance
err := DB.Where("provider_id = ?", providerID).Find(&instances).Error
if err != nil {
return nil, err
}
return instances, nil
}

func (dao *TenantModelInstanceDAO) GetByProviderIDAndInstanceName(providerID, instanceName string) (*entity.TenantModelInstance, error) {
var instance entity.TenantModelInstance
err := DB.Where("provider_id = ? AND instance_name = ?", providerID, instanceName).First(&instance).Error
Expand All @@ -64,3 +55,13 @@ func (dao *TenantModelInstanceDAO) DeleteByProviderIDAndInstanceName(providerID,
result := DB.Unscoped().Where("provider_id = ? and instance_name = ?", providerID, instanceName).Delete(&entity.TenantModelInstance{})
return result.RowsAffected, result.Error
}

// GetAllInstancesByProviderIDs get all instances by provider IDs
func (dao *TenantModelInstanceDAO) GetAllInstancesByProviderIDs(providerIDs []string) ([]*entity.TenantModelInstance, error) {
var instances []*entity.TenantModelInstance
err := DB.Where("provider_id IN ?", providerIDs).Find(&instances).Error
if err != nil {
return nil, err
}
return instances, nil
}
15 changes: 8 additions & 7 deletions internal/dao/tenant_model_provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,11 +64,12 @@ func (dao *TenantModelProviderDAO) DeleteByTenantIDAndProviderName(tenantID, pro
return result.RowsAffected, result.Error
}

// ListByID list tenant model providers by ID
func (dao *TenantModelProviderDAO) ListByID(id string) ([]string, error) {
var providerNames []string
err := DB.Model(&entity.TenantModelProvider{}).
Where("tenant_id = ?", id).
Pluck("provider_name", &providerNames).Error
return providerNames, err
// ListByID list tenant model providers by tenant ID, returns full provider objects
func (dao *TenantModelProviderDAO) ListByID(tenantID string) ([]*entity.TenantModelProvider, error) {
var providers []*entity.TenantModelProvider
err := DB.Where("tenant_id = ?", tenantID).Find(&providers).Error
if err != nil {
return nil, err
}
return providers, nil
}
25 changes: 22 additions & 3 deletions internal/entity/model.go
Original file line number Diff line number Diff line change
Expand Up @@ -271,6 +271,7 @@ func (pm *ProviderManager) ListProviders() ([]map[string]interface{}, error) {

modelTypeSet := make(map[string]struct{})
for _, model := range provider.Models {

for _, modelType := range model.ModelTypes {
modelTypeSet[modelType] = struct{}{}
}
Expand All @@ -287,6 +288,9 @@ func (pm *ProviderManager) ListProviders() ([]map[string]interface{}, error) {
"model_types": modelTypes,
"url_suffix": provider.URLSuffix,
}
if (len(modelTypes) == 0) {
continue
}
providers = append(providers, providerData)
}

Expand All @@ -305,10 +309,25 @@ func (pm *ProviderManager) GetProviderByName(providerName string) (map[string]in
return nil, fmt.Errorf("provider '%s' not found", providerName)
}

modelTypeSet := make(map[string]struct{})
for _, model := range provider.Models {
if len(model.ModelTypes) == 0 {
continue
}
for _, modelType := range model.ModelTypes {
modelTypeSet[modelType] = struct{}{}
}
}

var modelTypes []string
for modelType := range modelTypeSet {
modelTypes = append(modelTypes, modelType)
}

providerInfo := map[string]interface{}{
"name": provider.Name,
"base_url": provider.URL,
"total_models": len(provider.Models),
"name": provider.Name,
"url": provider.URL,
"model_types": modelTypes,
}

return providerInfo, nil
Expand Down
39 changes: 38 additions & 1 deletion internal/handler/tenant.go
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ func (h *TenantHandler) GetModels(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"code": common.CodeSuccess,
"message": "success",
"data": defaultModels,
"data": gin.H{"models": defaultModels},
})
}

Expand Down Expand Up @@ -119,6 +119,43 @@ func (h *TenantHandler) SetModels(c *gin.Context) {
})
}

// GetAddedModels lists all added models for the current user's tenant
// @Summary List Added Models
// @Description List all models added to the current user's tenant
// @Tags models
// @Accept json
// @Produce json
// @Security ApiKeyAuth
// @Param type query string false "Model type filter (chat, embedding, rerank, asr, vision, tts, ocr)"
// @Success 200 {object} map[string]interface{}
// @Router /api/v1/models [get]
func (h *TenantHandler) GetAddedModels(c *gin.Context) {
user, errorCode, errorMessage := GetUser(c)
if errorCode != common.CodeSuccess {
jsonError(c, errorCode, errorMessage)
return
}

// Get optional model type filter from query params
modelTypeFilter := c.Query("type")

addedModels, err := h.tenantService.ListTenantAddedModels(user.ID, modelTypeFilter)
if err != nil {
c.JSON(http.StatusOK, gin.H{
"code": common.CodeExceptionError,
"message": err.Error(),
"data": nil,
})
return
}

c.JSON(http.StatusOK, gin.H{
"code": common.CodeSuccess,
"message": "success",
"data": addedModels,
})
}

// TenantInfo get tenant information
// @Summary Get Tenant Information
// @Description Get current user's tenant information (owner tenant)
Expand Down
7 changes: 4 additions & 3 deletions internal/router/router.go
Original file line number Diff line number Diff line change
Expand Up @@ -287,7 +287,7 @@ func (r *Router) Setup(engine *gin.Engine) {
// provider pool route group
provider := v1.Group("/providers")
{
provider.GET("/", r.providerHandler.ListProviders)
provider.GET("", r.providerHandler.ListProviders)
provider.PUT("/", r.providerHandler.AddProvider)
provider.GET("/:provider_name", r.providerHandler.ShowProvider)
provider.DELETE("/:provider_name", r.providerHandler.DeleteProvider)
Expand Down Expand Up @@ -317,8 +317,9 @@ func (r *Router) Setup(engine *gin.Engine) {

model := v1.Group("/models")
{
model.GET("/", r.tenantHandler.GetModels)
model.PATCH("/", r.tenantHandler.SetModels)
model.GET("/default", r.tenantHandler.GetModels)
model.PATCH("/default", r.tenantHandler.SetModels)
model.GET("", r.tenantHandler.GetAddedModels) // GET /api/v1/models - list tenant added models
}

// Agent routes
Expand Down
38 changes: 20 additions & 18 deletions internal/service/model_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -131,14 +131,14 @@ func (m *ModelProviderService) ListProvidersOfTenant(userID string) ([]map[strin

tenantID := tenants[0].TenantID

providerNames, err := m.modelProviderDAO.ListByID(tenantID)
providers, err := m.modelProviderDAO.ListByID(tenantID)
if err != nil {
return nil, common.CodeServerError, err
}

var result []map[string]interface{}
for _, providerName := range providerNames {
provider, err := dao.GetModelProviderManager().GetProviderByName(providerName)
for _, p := range providers {
provider, err := dao.GetModelProviderManager().GetProviderByName(p.ProviderName)
if err != nil {
return nil, common.CodeServerError, err
}
Expand Down Expand Up @@ -295,7 +295,7 @@ func (m *ModelProviderService) ListProviderInstances(providerName, userID string
}

// Check if provider exists
instances, err := m.modelInstanceDAO.GetAllInstancesByProviderID(provider.ID)
instances, err := m.modelInstanceDAO.GetAllInstancesByProviderIDs([]string{provider.ID})
if err != nil {
return nil, common.CodeServerError, err
}
Expand All @@ -306,16 +306,16 @@ func (m *ModelProviderService) ListProviderInstances(providerName, userID string
var extra map[string]string
err = json.Unmarshal([]byte(instance.Extra), &extra)
if err != nil {
return nil, common.CodeServerError, err
extra = make(map[string]string)
}

result = append(result, map[string]interface{}{
"id": instance.ID,
"instanceName": instance.InstanceName,
"providerID": instance.ProviderID,
"apiKey": instance.APIKey,
"status": instance.Status,
"extra": instance.Extra,
"api_key": instance.APIKey,
"id": instance.ID,
"instance_name": instance.InstanceName,
"provider_id": instance.ProviderID,
"region": extra["region"],
"status": instance.Status,
})
}

Expand Down Expand Up @@ -351,15 +351,15 @@ func (m *ModelProviderService) ShowProviderInstance(providerName, instanceName,
var extra map[string]string
err = json.Unmarshal([]byte(instance.Extra), &extra)
if err != nil {
return nil, common.CodeServerError, err
extra = make(map[string]string)
}

result := map[string]interface{}{
"id": instance.ID,
"instanceName": instance.InstanceName,
"providerID": instance.ProviderID,
"status": instance.Status,
"region": extra["region"],
"id": instance.ID,
"instance_name": instance.InstanceName,
"provider_id": instance.ProviderID,
"status": instance.Status,
"region": extra["region"],
}

return result, common.CodeSuccess, nil
Expand Down Expand Up @@ -720,7 +720,7 @@ func (m *ModelProviderService) ListInstanceModels(providerName, instanceName, us
}

// Get all models for this instance
disabledModels, err := m.modelDAO.GetModelsByInstanceID(instance.ID)
disabledModels, err := m.modelDAO.GetModelsByInstanceIDs([]string{instance.ID})
if err != nil {
return nil, err
}
Expand All @@ -744,6 +744,8 @@ func (m *ModelProviderService) ListInstanceModels(providerName, instanceName, us
for _, model := range allModels {
// convert model["name"] to string
modelName := model["name"].(string)
model["model_type"] = model["model_types"]
delete(model, "model_types")
if modelNames[modelName] {
model["status"] = "inactive"
} else {
Expand Down
Loading