From 8b4673fe84a286f1bf9d17d28c6a4234e7dccee1 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 11 Jul 2025 22:24:31 -0700 Subject: [PATCH] build(ui/): UI - Public Model Hub v2 (#12532) Closes https://github.com/BerriAI/litellm/pull/12532 --- .../litellm_enterprise/proxy/proxy_server.py | 12 + litellm/constants.py | 96 +- .../_experimental/out/model_hub_table.html | 1 - .../proxy/_experimental/out/onboarding.html | 1 - .../ui_discovery_endpoints.py | 4 +- .../model_management_endpoints.py | 93 +- litellm/proxy/proxy_server.py | 34 +- .../public_endpoints/public_endpoints.py | 33 +- .../model_management_endpoints.py | 7 + .../public_endpoints/public_endpoints.py | 10 + .../src/components/model_hub_table.tsx | 132 ++- .../components/model_hub_table_columns.tsx | 28 +- .../src/components/networking.tsx | 43 + .../src/components/public_model_hub.tsx | 870 ++++++++++++++++++ .../components/public_model_hub_columns.tsx | 218 +++++ 15 files changed, 1522 insertions(+), 60 deletions(-) delete mode 100644 litellm/proxy/_experimental/out/model_hub_table.html delete mode 100644 litellm/proxy/_experimental/out/onboarding.html create mode 100644 litellm/types/proxy/management_endpoints/model_management_endpoints.py create mode 100644 litellm/types/proxy/public_endpoints/public_endpoints.py create mode 100644 ui/litellm-dashboard/src/components/public_model_hub.tsx create mode 100644 ui/litellm-dashboard/src/components/public_model_hub_columns.tsx diff --git a/enterprise/litellm_enterprise/proxy/proxy_server.py b/enterprise/litellm_enterprise/proxy/proxy_server.py index 96503f172a1..79d3ebdf9ee 100644 --- a/enterprise/litellm_enterprise/proxy/proxy_server.py +++ b/enterprise/litellm_enterprise/proxy/proxy_server.py @@ -1,3 +1,4 @@ +import os from typing import Optional from litellm_enterprise.types.proxy.proxy_server import CustomAuthSettings @@ -20,3 +21,14 @@ class EnterpriseProxyConfig: global custom_auth_settings custom_auth_settings = await self.load_custom_auth_settings(general_settings) return None + + @staticmethod + def get_custom_docs_description() -> Optional[str]: + from litellm.proxy.proxy_server import premium_user + + docs_description: Optional[str] = None + if premium_user: + # check if premium_user has custom_docs_description + docs_description = os.getenv("DOCS_DESCRIPTION") + + return docs_description diff --git a/litellm/constants.py b/litellm/constants.py index 67e2a1589c9..733af4ef616 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -113,16 +113,24 @@ AZURE_FILE_SEARCH_COST_PER_GB_PER_DAY = float( os.getenv("AZURE_FILE_SEARCH_COST_PER_GB_PER_DAY", 0.1) # $0.1 USD per 1 GB/Day ) AZURE_CODE_INTERPRETER_COST_PER_SESSION = float( - os.getenv("AZURE_CODE_INTERPRETER_COST_PER_SESSION", 0.03) # $0.03 USD per 1 Session + os.getenv( + "AZURE_CODE_INTERPRETER_COST_PER_SESSION", 0.03 + ) # $0.03 USD per 1 Session ) AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS = float( - os.getenv("AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS", 3.0) # $0.003 USD per 1K Tokens + os.getenv( + "AZURE_COMPUTER_USE_INPUT_COST_PER_1K_TOKENS", 3.0 + ) # $0.003 USD per 1K Tokens ) AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS = float( - os.getenv("AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS", 12.0) # $0.012 USD per 1K Tokens + os.getenv( + "AZURE_COMPUTER_USE_OUTPUT_COST_PER_1K_TOKENS", 12.0 + ) # $0.012 USD per 1K Tokens ) AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY = float( - os.getenv("AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY", 0.1) # $0.1 USD per 1 GB/Day (same as file search) + os.getenv( + "AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY", 0.1 + ) # $0.1 USD per 1 GB/Day (same as file search) ) MIN_NON_ZERO_TEMPERATURE = float(os.getenv("MIN_NON_ZERO_TEMPERATURE", 0.0001)) #### RELIABILITY #### @@ -394,7 +402,7 @@ openai_compatible_endpoints: List = [ "api.featherless.ai/v1", "inference.api.nscale.com/v1", "api.studio.nebius.ai/v1", - "https://dashscope-intl.aliyuncs.com/compatible-mode/v1" + "https://dashscope-intl.aliyuncs.com/compatible-mode/v1", ] @@ -430,7 +438,7 @@ openai_compatible_providers: List = [ "featherless_ai", "nscale", "nebius", - "dashscope" + "dashscope", ] openai_text_completion_compatible_providers: List = ( [ # providers that support `/v1/completions` @@ -441,7 +449,7 @@ openai_text_completion_compatible_providers: List = ( "llamafile", "featherless_ai", "nebius", - "dashscope" + "dashscope", ] ) _openai_like_providers: List = [ @@ -625,7 +633,7 @@ dashscope_models: List = [ "qwq-32b", "qwen3-235b-a22b", "qwen3-32b", - "qwen3-30b-a3b" + "qwen3-30b-a3b", ] nebius_embedding_models: List = [ @@ -816,7 +824,10 @@ LENGTH_OF_LITELLM_GENERATED_KEY = int(os.getenv("LENGTH_OF_LITELLM_GENERATED_KEY SECRET_MANAGER_REFRESH_INTERVAL = int( os.getenv("SECRET_MANAGER_REFRESH_INTERVAL", 86400) ) -LITELLM_SETTINGS_SAFE_DB_OVERRIDES = ["default_internal_user_params"] +LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [ + "default_internal_user_params", + "public_model_groups", +] SPECIAL_LITELLM_AUTH_TOKEN = ["ui-token"] DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int( os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60) @@ -825,25 +836,62 @@ DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int( # Sentry Scrubbing Configuration SENTRY_DENYLIST = [ # API Keys and Tokens - "api_key", "token", "key", "secret", "password", "auth", "credential", - "OPENAI_API_KEY", "ANTHROPIC_API_KEY", "AZURE_API_KEY", "COHERE_API_KEY", - "REPLICATE_API_KEY", "HUGGINGFACE_API_KEY", "TOGETHERAI_API_KEY", - "CLOUDFLARE_API_KEY", "BASETEN_KEY", "OPENROUTER_KEY", "DATAROBOT_API_TOKEN", - "FIREWORKS_API_KEY", "FIREWORKS_AI_API_KEY", "FIREWORKSAI_API_KEY", + "api_key", + "token", + "key", + "secret", + "password", + "auth", + "credential", + "OPENAI_API_KEY", + "ANTHROPIC_API_KEY", + "AZURE_API_KEY", + "COHERE_API_KEY", + "REPLICATE_API_KEY", + "HUGGINGFACE_API_KEY", + "TOGETHERAI_API_KEY", + "CLOUDFLARE_API_KEY", + "BASETEN_KEY", + "OPENROUTER_KEY", + "DATAROBOT_API_TOKEN", + "FIREWORKS_API_KEY", + "FIREWORKS_AI_API_KEY", + "FIREWORKSAI_API_KEY", # Database and Connection Strings - "database_url", "redis_url", "connection_string", + "database_url", + "redis_url", + "connection_string", # Authentication and Security - "master_key", "LITELLM_MASTER_KEY", "auth_token", "jwt_token", "private_key", - "SLACK_WEBHOOK_URL", "webhook_url", "LANGFUSE_SECRET_KEY", + "master_key", + "LITELLM_MASTER_KEY", + "auth_token", + "jwt_token", + "private_key", + "SLACK_WEBHOOK_URL", + "webhook_url", + "LANGFUSE_SECRET_KEY", # Email Configuration - "SMTP_PASSWORD", "SMTP_USERNAME", "email_password", + "SMTP_PASSWORD", + "SMTP_USERNAME", + "email_password", # Cloud Provider Credentials - "aws_access_key", "aws_secret_key", "gcp_credentials", - "azure_credentials", "HCP_VAULT_TOKEN", "CIRCLE_OIDC_TOKEN", + "aws_access_key", + "aws_secret_key", + "gcp_credentials", + "azure_credentials", + "HCP_VAULT_TOKEN", + "CIRCLE_OIDC_TOKEN", # Proxy and Environment Settings - "proxy_url", "proxy_key", "environment_variables" + "proxy_url", + "proxy_key", + "environment_variables", ] SENTRY_PII_DENYLIST = [ - "user_id", "email", "phone", "address", "ip_address", - "SMTP_SENDER_EMAIL", "TEST_EMAIL_ADDRESS" -] \ No newline at end of file + "user_id", + "email", + "phone", + "address", + "ip_address", + "SMTP_SENDER_EMAIL", + "TEST_EMAIL_ADDRESS", +] diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table.html deleted file mode 100644 index 5216e82518b..00000000000 --- a/litellm/proxy/_experimental/out/model_hub_table.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index 1f48cbbb4c8..00000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py b/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py index 2a6bec77bfb..b5749046897 100644 --- a/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py +++ b/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py @@ -13,8 +13,10 @@ router = APIRouter() "/litellm/.well-known/litellm-ui-config", response_model=UiDiscoveryEndpoints ) # if mounted at root path async def get_ui_config(): + from litellm.proxy.proxy_server import _title, version from litellm.proxy.utils import get_proxy_base_url, get_server_root_path return UiDiscoveryEndpoints( - server_root_path=get_server_root_path(), proxy_base_url=get_proxy_base_url() + server_root_path=get_server_root_path(), + proxy_base_url=get_proxy_base_url(), ) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index e4189c6ce3d..8330f403816 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -17,7 +17,7 @@ import uuid from typing import Dict, List, Literal, Optional, Tuple, Union, cast from fastapi import APIRouter, Depends, HTTPException, Request, status -from pydantic import BaseModel +from pydantic import BaseModel, ConfigDict, Field from litellm._logging import verbose_proxy_logger from litellm.constants import LITELLM_PROXY_ADMIN_NAME @@ -48,6 +48,7 @@ from litellm.types.router import ( Deployment, DeploymentTypedDict, LiteLLMParamsTypedDict, + ModelGroupInfo, updateDeployment, ) from litellm.utils import get_utc_datetime @@ -55,6 +56,16 @@ from litellm.utils import get_utc_datetime router = APIRouter() +class UpdatePublicModelGroupsRequest(BaseModel): + """Request model for updating public model groups""" + + model_groups: List[str] = Field( + description="List of model group names to make public" + ) + + model_config = ConfigDict(extra="forbid") + + async def get_db_model( model_id: str, prisma_client: PrismaClient ) -> Optional[Deployment]: @@ -945,6 +956,86 @@ async def update_model( ) +@router.post( + "/model_group/make_public", + description="Update which model groups are public", + tags=["model management"], + dependencies=[Depends(user_api_key_auth)], +) +async def update_public_model_groups( + request: UpdatePublicModelGroupsRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Update which model groups are public. + + This endpoint allows admins to specify which model groups should be publicly accessible. + Public model groups are visible via the /public/model_hub endpoint. + + Args: + request: Request containing list of model group names to make public + user_api_key_dict: User authentication information + + Returns: + Success message with updated public model groups + + Raises: + ProxyException: For various error conditions including authentication errors + """ + try: + # Update the public model groups + import litellm + from litellm.proxy.proxy_server import proxy_config + + # Check if user has admin permissions + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail={ + "error": "Only proxy admins can update public model groups. Your role={}".format( + user_api_key_dict.user_role + ) + }, + ) + + litellm.public_model_groups = request.model_groups + + # Load existing config + config = await proxy_config.get_config() + + # Update config with new settings + if "litellm_settings" not in config: + config["litellm_settings"] = {} + + config["litellm_settings"]["public_model_groups"] = request.model_groups + + # Save the updated config + await proxy_config.save_config(new_config=config) + + verbose_proxy_logger.info( + f"Updated public model groups to: {request.model_groups} by user: {user_api_key_dict.user_id}" + ) + + return { + "message": "Successfully updated public model groups", + "public_model_groups": request.model_groups, + "updated_by": user_api_key_dict.user_id, + } + + except Exception as e: + verbose_proxy_logger.exception(f"Error updating public model groups: {str(e)}") + + if isinstance(e, HTTPException): + raise e + + raise ProxyException( + message=f"Error updating public model groups: {str(e)}", + type=ProxyErrorTypes.internal_server_error, + code=status.HTTP_500_INTERNAL_SERVER_ERROR, + param=None, + ) + + def _deduplicate_litellm_router_models(models: List[Dict]) -> List[Dict]: """ Deduplicate models based on their model_info.id field. diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ea2222092ba..6d250bb8ae2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -355,6 +355,9 @@ from litellm.types.llms.anthropic import ( AnthropicResponseUsageBlock, ) from litellm.types.llms.openai import HttpxBinaryResponseContent +from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ModelGroupInfoProxy, +) from litellm.types.proxy.management_endpoints.ui_sso import ( DefaultTeamSSOParams, LiteLLM_UpperboundKeyGenerateParams, @@ -1640,7 +1643,7 @@ class ProxyConfig: if credential_list_dict: credential_list = [CredentialItem(**cred) for cred in credential_list_dict] return credential_list - + def _load_environment_variables(self, config: dict): ## ENVIRONMENT VARIABLES global premium_user @@ -1655,7 +1658,9 @@ class ProxyConfig: # ``` ######################################################### if isinstance(value, str) and value.startswith("os.environ/"): - resolved_secret_string: Optional[str] = get_secret_str(secret_name=value) + resolved_secret_string: Optional[str] = get_secret_str( + secret_name=value + ) if resolved_secret_string is not None: os.environ[key] = resolved_secret_string else: @@ -1974,10 +1979,13 @@ class ProxyConfig: value=custom_sso, config_file_path=config_file_path ) - custom_ui_sso_sign_in_handler = general_settings.get("custom_ui_sso_sign_in_handler", None) + custom_ui_sso_sign_in_handler = general_settings.get( + "custom_ui_sso_sign_in_handler", None + ) if custom_ui_sso_sign_in_handler is not None: user_custom_ui_sso_sign_in_handler = get_instance_fn( - value=custom_ui_sso_sign_in_handler, config_file_path=config_file_path + value=custom_ui_sso_sign_in_handler, + config_file_path=config_file_path, ) if enterprise_proxy_config is not None: @@ -6530,8 +6538,8 @@ async def model_info_v1( # noqa: PLR0915 def _get_model_group_info( llm_router: Router, all_models_str: List[str], model_group: Optional[str] -) -> List[ModelGroupInfo]: - model_groups: List[ModelGroupInfo] = [] +) -> List[ModelGroupInfoProxy]: + model_groups: List[ModelGroupInfoProxy] = [] # ensure all_models_str is a set all_models_str_set = set(all_models_str) @@ -6542,15 +6550,23 @@ def _get_model_group_info( _model_group_info = llm_router.get_model_group_info(model_group=model) if _model_group_info is not None: - model_groups.append(_model_group_info) + model_groups.append(ModelGroupInfoProxy(**_model_group_info.model_dump())) else: - model_group_info = ModelGroupInfo( + model_group_info = ModelGroupInfoProxy( model_group=model, providers=[], ) model_groups.append(model_group_info) + + ## check for public model groups + if litellm.public_model_groups is not None: + for mg in model_groups: + if mg.model_group in litellm.public_model_groups: + mg.is_public_model_group = True + return model_groups + @router.get( "/model_group/info", tags=["model management"], @@ -6772,7 +6788,7 @@ async def model_group_info( infer_model_from_keys=general_settings.get("infer_model_from_keys", False), llm_router=llm_router, ) - model_groups: List[ModelGroupInfo] = _get_model_group_info( + model_groups: List[ModelGroupInfoProxy] = _get_model_group_info( llm_router=llm_router, all_models_str=all_models_str, model_group=model_group ) diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index af33d68905e..4910f71429e 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -4,7 +4,10 @@ from fastapi import APIRouter, Depends, HTTPException from litellm.proxy._types import CommonProxyErrors from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.types.router import ModelGroupInfo +from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ModelGroupInfoProxy, +) +from litellm.types.proxy.public_endpoints.public_endpoints import PublicModelHubInfo router = APIRouter() @@ -13,7 +16,7 @@ router = APIRouter() "/public/model_hub", tags=["public", "model management"], dependencies=[Depends(user_api_key_auth)], - response_model=List[ModelGroupInfo], + response_model=List[ModelGroupInfoProxy], ) async def public_model_hub(): import litellm @@ -24,7 +27,7 @@ async def public_model_hub(): status_code=400, detail=CommonProxyErrors.no_llm_router.value ) - model_groups: List[ModelGroupInfo] = [] + model_groups: List[ModelGroupInfoProxy] = [] if litellm.public_model_groups is not None: model_groups = _get_model_group_info( llm_router=llm_router, @@ -33,3 +36,27 @@ async def public_model_hub(): ) return model_groups + + +@router.get( + "/public/model_hub/info", + tags=["public", "model management"], + response_model=PublicModelHubInfo, +) +async def public_model_hub_info(): + import litellm + from litellm.proxy.proxy_server import _title, version + + try: + from litellm_enterprise.proxy.proxy_server import EnterpriseProxyConfig + + custom_docs_description = EnterpriseProxyConfig.get_custom_docs_description() + except Exception: + custom_docs_description = None + + return PublicModelHubInfo( + docs_title=_title, + custom_docs_description=custom_docs_description, + litellm_version=version, + useful_links=litellm.public_model_groups_links, + ) diff --git a/litellm/types/proxy/management_endpoints/model_management_endpoints.py b/litellm/types/proxy/management_endpoints/model_management_endpoints.py new file mode 100644 index 00000000000..16403b663be --- /dev/null +++ b/litellm/types/proxy/management_endpoints/model_management_endpoints.py @@ -0,0 +1,7 @@ +from pydantic import Field + +from ...router import ModelGroupInfo + + +class ModelGroupInfoProxy(ModelGroupInfo): + is_public_model_group: bool = Field(default=False) diff --git a/litellm/types/proxy/public_endpoints/public_endpoints.py b/litellm/types/proxy/public_endpoints/public_endpoints.py new file mode 100644 index 00000000000..b2949a719ed --- /dev/null +++ b/litellm/types/proxy/public_endpoints/public_endpoints.py @@ -0,0 +1,10 @@ +from typing import Dict, Optional + +from pydantic import BaseModel + + +class PublicModelHubInfo(BaseModel): + docs_title: str + custom_docs_description: Optional[str] + litellm_version: str + useful_links: Optional[Dict[str, str]] diff --git a/ui/litellm-dashboard/src/components/model_hub_table.tsx b/ui/litellm-dashboard/src/components/model_hub_table.tsx index 4a5746b9b56..c0a269c38f1 100644 --- a/ui/litellm-dashboard/src/components/model_hub_table.tsx +++ b/ui/litellm-dashboard/src/components/model_hub_table.tsx @@ -1,9 +1,10 @@ import React, { useEffect, useState, useRef } from "react"; import { useRouter, useSearchParams } from "next/navigation"; -import { modelHubCall } from "./networking"; +import { modelHubCall, makeModelGroupPublic, modelHubPublicModelsCall, proxyBaseUrl } from "./networking"; import { getConfigFieldSetting, updateConfigFieldSetting } from "./networking"; import { ModelDataTable } from "./model_dashboard/table"; import { modelHubColumns } from "./model_hub_table_columns"; +import PublicModelHub from "./public_model_hub"; import { Card, Text, @@ -55,16 +56,13 @@ const ModelHubTable: React.FC = ({ const [searchTerm, setSearchTerm] = useState(""); const [selectedProvider, setSelectedProvider] = useState(""); const [selectedMode, setSelectedMode] = useState(""); + const [selectedFeature, setSelectedFeature] = useState(""); const [selectedModels, setSelectedModels] = useState>(new Set()); const router = useRouter(); const tableRef = useRef>(null); useEffect(() => { - if (!accessToken) { - return; - } - - const fetchData = async () => { + const fetchData = async (accessToken: string) => { try { setLoading(true); const _modelHubData = await modelHubCall(accessToken); @@ -88,7 +86,28 @@ const ModelHubTable: React.FC = ({ } }; - fetchData(); + const fetchPublicData = async () => { + try { + setLoading(true); + const _modelHubData = await modelHubPublicModelsCall(); + console.log("ModelHubData:", _modelHubData); + console.log("First model structure:", _modelHubData[0]); + console.log("Model has model_group?", _modelHubData[0]?.model_group); + console.log("Model has providers?", _modelHubData[0]?.providers); + setModelHubData(_modelHubData); + setPublicPageAllowed(true); + } catch (error) { + console.error("There was an error fetching the public model data", error); + } finally { + setLoading(false); + } + } + + if (accessToken) { + fetchData(accessToken); + } else if (publicPage) { + fetchPublicData(); + } }, [accessToken, publicPage]); const showModal = (model: ModelGroupInfo) => { @@ -104,11 +123,30 @@ const ModelHubTable: React.FC = ({ if (!accessToken) { return; } - updateConfigFieldSetting(accessToken, "enable_public_model_hub", true).then( - (data) => { + + try { + // Get the selected model groups or use all if none are selected + const modelGroupsToMakePublic = selectedModels.size > 0 + ? Array.from(selectedModels) + : modelHubData?.map(model => model.model_group) || []; + + if (modelGroupsToMakePublic.length > 0) { + // Call the endpoint to make the selected model groups public + await makeModelGroupPublic(accessToken, modelGroupsToMakePublic); + + // Show success message + message.success(`Successfully made ${modelGroupsToMakePublic.length} model group(s) public!`); + + // Route to the model hub table + router.push(`/ui/model_hub_table`); + } else { + // Show the modal if no model groups available setIsPublicPageModalVisible(true); } - ); + } catch (error) { + console.error("Error making model groups public:", error); + message.error("Failed to make model groups public. Please try again."); + } }; const handleOk = () => { @@ -164,11 +202,44 @@ const ModelHubTable: React.FC = ({ return Array.from(modes); }; + const getUniqueFeatures = (data: ModelGroupInfo[]) => { + const features = new Set(); + data.forEach(model => { + // Find all properties that start with 'supports_' and are true + Object.entries(model) + .filter(([key, value]) => key.startsWith('supports_') && value === true) + .forEach(([key]) => { + // Format the feature name (remove 'supports_' prefix and convert to title case) + const featureName = key + .replace(/^supports_/, '') + .split('_') + .map(word => word.charAt(0).toUpperCase() + word.slice(1)) + .join(' '); + features.add(featureName); + }); + }); + return Array.from(features).sort(); + }; + const filteredData = modelHubData?.filter(model => { const matchesSearch = model.model_group.toLowerCase().includes(searchTerm.toLowerCase()); const matchesProvider = selectedProvider === "" || model.providers.includes(selectedProvider); const matchesMode = selectedMode === "" || model.mode === selectedMode; - return matchesSearch && matchesProvider && matchesMode; + + // Check if model has the selected feature + const matchesFeature = selectedFeature === "" || + Object.entries(model) + .filter(([key, value]) => key.startsWith('supports_') && value === true) + .some(([key]) => { + const featureName = key + .replace(/^supports_/, '') + .split('_') + .map(word => word.charAt(0).toUpperCase() + word.slice(1)) + .join(' '); + return featureName === selectedFeature; + }); + + return matchesSearch && matchesProvider && matchesMode && matchesFeature; }) || []; const handleRowSelection = (modelGroup: string, isSelected: boolean) => { @@ -198,17 +269,33 @@ const ModelHubTable: React.FC = ({ // Clear selections when filters change to avoid confusion useEffect(() => { setSelectedModels(new Set()); - }, [searchTerm, selectedProvider, selectedMode]); + }, [searchTerm, selectedProvider, selectedMode, selectedFeature]); + console.log("publicPage: ", publicPage); + console.log("publicPageAllowed: ", publicPageAllowed); + + // If this is a public page, use the dedicated PublicModelHub component + if (publicPage && publicPageAllowed) { + return ; + } + return (
- {(publicPage && publicPageAllowed) || publicPage == false ? ( + {publicPage == false ? (
Model Hub - Table View +
+ Model Hub URL: + {`${proxyBaseUrl}/ui/model_hub_table`} + {publicPage == false ? ( premiumUser ? ( - ) : ( @@ -225,6 +312,7 @@ const ModelHubTable: React.FC = ({
)}
+
{/* Filters */} @@ -265,6 +353,19 @@ const ModelHubTable: React.FC = ({ ))}
+
+ Features: + +
@@ -278,6 +379,7 @@ const ModelHubTable: React.FC = ({ handleSelectAll, showModal, copyToClipboard, + publicPage, )} data={filteredData} isLoading={loading} @@ -330,7 +432,7 @@ const ModelHubTable: React.FC = ({
Shareable Link: - {`/ui/model_hub_table?key=`} + {`${proxyBaseUrl}/model_hub_table`}
diff --git a/ui/litellm-dashboard/src/components/model_hub_table_columns.tsx b/ui/litellm-dashboard/src/components/model_hub_table_columns.tsx index b25916497b9..76e443a6910 100644 --- a/ui/litellm-dashboard/src/components/model_hub_table_columns.tsx +++ b/ui/litellm-dashboard/src/components/model_hub_table_columns.tsx @@ -59,7 +59,9 @@ export const modelHubColumns = ( handleSelectAll: (checked: boolean) => void, showModal: (model: ModelHubData) => void, copyToClipboard: (text: string) => void, -): ColumnDef[] => [ + publicPage: boolean = false, +): ColumnDef[] => { + const allColumns: ColumnDef[] = [ { header: () => ( { - const publicA = rowA.original.public === true ? 1 : 0; - const publicB = rowB.original.public === true ? 1 : 0; + const publicA = rowA.original.is_public_model_group === true ? 1 : 0; + const publicB = rowB.original.is_public_model_group === true ? 1 : 0; return publicA - publicB; }, cell: ({ row }) => { const model = row.original; - return model.public === true ? ( + return model.is_public_model_group === true ? ( Yes ) : ( No @@ -278,4 +280,20 @@ export const modelHubColumns = ( ); }, }, -]; \ No newline at end of file +]; + + // Filter out columns based on publicPage setting + if (publicPage) { + return allColumns.filter(column => { + // Remove the select/checkbox column + if (column.id === "select") return false; + + // Remove the public column + if ('accessorKey' in column && column.accessorKey === "public") return false; + + return true; + }); + } + + return allColumns; +}; \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 1a34eed5833..97940d19eb7 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -108,6 +108,13 @@ export interface CredentialItem { }; } +export interface PublicModelHubInfo { + docs_title: string; + custom_docs_description: string | null; + litellm_version: string; + useful_links: Record; +} + export interface LiteLLMWellKnownUiConfig { server_root_path: string; proxy_base_url: string | null; @@ -148,6 +155,21 @@ export function setGlobalLitellmHeaderName( globalLitellmHeaderName = headerName; } +export const makeModelGroupPublic = async (accessToken: string, modelGroups: string[]) => { + const url = proxyBaseUrl ? `${proxyBaseUrl}/model_group/make_public` : `/model_group/make_public`; + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify({ + model_groups: modelGroups, + }), + }); + return response.json(); +}; + export const getUiConfig = async () => { console.log("Getting UI config"); /**Special route to get the proxy base url and server root path */ @@ -164,6 +186,15 @@ export const getUiConfig = async () => { return jsonData; }; +export const getPublicModelHubInfo = async () => { + const url = defaultProxyBaseUrl + ? `${defaultProxyBaseUrl}/public/model_hub/info` + : `/public/model_hub/info`; + const response = await fetch(url); + const jsonData: PublicModelHubInfo = await response.json(); + return jsonData; +}; + export const getOpenAPISchema = async () => { const url = proxyBaseUrl ? `${proxyBaseUrl}/openapi.json` : `/openapi.json`; const response = await fetch(url); @@ -1769,6 +1800,17 @@ export const modelInfoV1Call = async (accessToken: String, modelId: String) => { } }; +export const modelHubPublicModelsCall = async () => { + const url = proxyBaseUrl ? `${proxyBaseUrl}/public/model_hub` : `/public/model_hub`; + const response = await fetch(url, { + method: "GET", + headers: { + "Content-Type": "application/json", + }, + }); + return response.json(); +}; + export const modelHubCall = async (accessToken: String) => { /** * Get all models on proxy @@ -4577,6 +4619,7 @@ export const healthCheckHistoryCall = async ( } }; + export const latestHealthChecksCall = async (accessToken: String) => { /** * Get the latest health check status for all models diff --git a/ui/litellm-dashboard/src/components/public_model_hub.tsx b/ui/litellm-dashboard/src/components/public_model_hub.tsx new file mode 100644 index 00000000000..5625ff3377e --- /dev/null +++ b/ui/litellm-dashboard/src/components/public_model_hub.tsx @@ -0,0 +1,870 @@ +import React, { useEffect, useState, useRef, useMemo } from "react"; +import { modelHubPublicModelsCall, proxyBaseUrl, getUiConfig, getPublicModelHubInfo } from "./networking"; +import { ModelDataTable } from "./model_dashboard/table"; +import { ColumnDef } from "@tanstack/react-table"; +import { + Card, + Text, + Title, + Button, +} from "@tremor/react"; +import { message, Tag, Tooltip, Modal, Select } from "antd"; +import { CopyOutlined } from "@ant-design/icons"; +import { ExternalLinkIcon, SearchIcon, EyeIcon, CogIcon } from "@heroicons/react/outline"; +import { Copy, Info } from "lucide-react"; +import { Table as TableInstance } from '@tanstack/react-table'; +import { generateCodeSnippet } from "./chat_ui/CodeSnippets"; +import { EndpointType, getEndpointType } from "./chat_ui/mode_endpoint_mapping"; +import { MessageType } from "./chat_ui/types"; +import { getProviderLogoAndName } from "./provider_info_helpers"; +// Simple approach without react-markdown dependency + +interface ModelGroupInfo { + model_group: string; + providers: string[]; + max_input_tokens?: number; + max_output_tokens?: number; + input_cost_per_token?: number; + output_cost_per_token?: number; + mode?: string; + tpm?: number; + rpm?: number; + supports_parallel_function_calling: boolean; + supports_vision: boolean; + supports_function_calling: boolean; + supported_openai_params?: string[]; + [key: string]: any; +} + +interface PublicModelHubProps { + accessToken?: string | null; +} + +const PublicModelHub: React.FC = ({ accessToken }) => { + const [modelHubData, setModelHubData] = useState(null); + const [pageTitle, setPageTitle] = useState("LiteLLM Gateway"); + const [customDocsDescription, setCustomDocsDescription] = useState(null); + const [litellmVersion, setLitellmVersion] = useState(""); + const [usefulLinks, setUsefulLinks] = useState>({}); + const [loading, setLoading] = useState(true); + const [searchTerm, setSearchTerm] = useState(""); + const [selectedProviders, setSelectedProviders] = useState([]); + const [selectedModes, setSelectedModes] = useState([]); + const [selectedFeatures, setSelectedFeatures] = useState([]); + const [serviceStatus, setServiceStatus] = useState("I'm alive! ✓"); + const [isModalVisible, setIsModalVisible] = useState(false); + const [selectedModel, setSelectedModel] = useState(null); + const tableRef = useRef>(null); + + useEffect(() => { + const fetchPublicData = async () => { + try { + setLoading(true); + const _modelHubData = await modelHubPublicModelsCall(); + console.log("ModelHubData:", _modelHubData); + setModelHubData(_modelHubData); + } catch (error) { + console.error("There was an error fetching the public model data", error); + setServiceStatus("Service unavailable"); + } finally { + setLoading(false); + } + }; + + const fetchPublicModelHubInfo = async () => { + const publicModelHubInfo = await getPublicModelHubInfo(); + console.log("Public Model Hub Info:", publicModelHubInfo); + setPageTitle(publicModelHubInfo.docs_title); + setCustomDocsDescription(publicModelHubInfo.custom_docs_description); + setLitellmVersion(publicModelHubInfo.litellm_version); + setUsefulLinks(publicModelHubInfo.useful_links || {}); + }; + + + fetchPublicModelHubInfo(); + + fetchPublicData(); + }, []); + + // Clear filters when filter values change to avoid confusion + useEffect(() => { + // This would clear selections if we had any selection functionality + // For now, it's just for consistency with the original component + }, [searchTerm, selectedProviders, selectedModes, selectedFeatures]); + + const getUniqueProviders = (data: ModelGroupInfo[]) => { + const providers = new Set(); + data.forEach(model => { + model.providers.forEach(provider => providers.add(provider)); + }); + return Array.from(providers); + }; + + const getUniqueModes = (data: ModelGroupInfo[]) => { + const modes = new Set(); + data.forEach(model => { + if (model.mode) modes.add(model.mode); + }); + return Array.from(modes); + }; + + const getUniqueFeatures = (data: ModelGroupInfo[]) => { + const features = new Set(); + data.forEach(model => { + // Find all properties that start with 'supports_' and are true + Object.entries(model) + .filter(([key, value]) => key.startsWith('supports_') && value === true) + .forEach(([key]) => { + // Format the feature name (remove 'supports_' prefix and convert to title case) + const featureName = key + .replace(/^supports_/, '') + .split('_') + .map(word => word.charAt(0).toUpperCase() + word.slice(1)) + .join(' '); + features.add(featureName); + }); + }); + return Array.from(features).sort(); + }; + + + + const filteredData = useMemo(() => { + if (!modelHubData) return []; + + let searchResults = modelHubData; + + // Apply search if there's a search term + if (searchTerm.trim()) { + const lowercaseSearch = searchTerm.toLowerCase(); + const searchWords = lowercaseSearch.split(/\s+/); + + // First, try flexible matching that handles different separators + const exactMatches = modelHubData.filter(model => { + const modelName = model.model_group.toLowerCase(); + + // Check if it contains the exact search term + if (modelName.includes(lowercaseSearch)) { + return true; + } + + // Check if it contains all search words (handles spaces vs slashes/dashes) + return searchWords.every(word => modelName.includes(word)); + }); + + // If we have exact matches, rank them by relevance + if (exactMatches.length > 0) { + searchResults = exactMatches.sort((a, b) => { + const aName = a.model_group.toLowerCase(); + const bName = b.model_group.toLowerCase(); + + // Calculate relevance scores + const aExactMatch = aName === lowercaseSearch ? 1000 : 0; + const bExactMatch = bName === lowercaseSearch ? 1000 : 0; + + const aStartsWith = aName.startsWith(lowercaseSearch) ? 100 : 0; + const bStartsWith = bName.startsWith(lowercaseSearch) ? 100 : 0; + + const aContainsWords = lowercaseSearch.split(/\s+/).every(word => aName.includes(word)) ? 50 : 0; + const bContainsWords = lowercaseSearch.split(/\s+/).every(word => bName.includes(word)) ? 50 : 0; + + const aLength = aName.length; + const bLength = bName.length; + + const aScore = aExactMatch + aStartsWith + aContainsWords + (1000 - aLength); + const bScore = bExactMatch + bStartsWith + bContainsWords + (1000 - bLength); + + return bScore - aScore; // Higher score first + }); + } + } + + // Apply other filters + return searchResults.filter(model => { + const matchesProvider = selectedProviders.length === 0 || selectedProviders.some(provider => model.providers.includes(provider)); + const matchesMode = selectedModes.length === 0 || selectedModes.includes(model.mode || ""); + + // Check if model has any of the selected features + const matchesFeature = selectedFeatures.length === 0 || + Object.entries(model) + .filter(([key, value]) => key.startsWith('supports_') && value === true) + .some(([key]) => { + const featureName = key + .replace(/^supports_/, '') + .split('_') + .map(word => word.charAt(0).toUpperCase() + word.slice(1)) + .join(' '); + return selectedFeatures.includes(featureName); + }); + + return matchesProvider && matchesMode && matchesFeature; + }); + }, [modelHubData, searchTerm, selectedProviders, selectedModes, selectedFeatures]); + + + + const showModal = (model: ModelGroupInfo) => { + setSelectedModel(model); + setIsModalVisible(true); + }; + + const handleModalOk = () => { + setIsModalVisible(false); + setSelectedModel(null); + }; + + const handleModalCancel = () => { + setIsModalVisible(false); + setSelectedModel(null); + }; + + const copyToClipboard = (text: string) => { + navigator.clipboard.writeText(text); + message.success("Copied to clipboard!"); + }; + + const formatCapabilityName = (key: string) => { + return key + .replace(/^supports_/, '') + .split('_') + .map(word => word.charAt(0).toUpperCase() + word.slice(1)) + .join(' '); + }; + + const getModelCapabilities = (model: ModelGroupInfo) => { + return Object.entries(model) + .filter(([key, value]) => key.startsWith('supports_') && value === true) + .map(([key]) => key); + }; + + + + const formatCost = (cost: number) => { + return `$${(cost * 1_000_000).toFixed(4)}`; + }; + + const formatTokens = (tokens: number | undefined) => { + if (!tokens) return "N/A"; + if (tokens >= 1000) { + return `${(tokens / 1000).toFixed(0)}K`; + } + return tokens.toString(); + }; + + const formatLimits = (rpm?: number, tpm?: number) => { + const limits = []; + if (rpm) limits.push(`RPM: ${rpm.toLocaleString()}`); + if (tpm) limits.push(`TPM: ${tpm.toLocaleString()}`); + return limits.length > 0 ? limits.join(", ") : "N/A"; + }; + + + + const publicModelHubColumns = (): ColumnDef[] => [ + { + header: "Model Name", + accessorKey: "model_group", + enableSorting: true, + cell: ({ row }) => ( +
+ + + +
+ ), + size: 150, + }, + { + header: "Providers", + accessorKey: "providers", + enableSorting: true, + cell: ({ row }) => { + const providers = row.original.providers; + + return ( +
+ {providers.map((provider) => { + const { logo } = getProviderLogoAndName(provider); + return ( +
+ {logo && ( + {provider} { + (e.target as HTMLImageElement).style.display = 'none'; + }} + /> + )} + {provider} +
+ ); + })} +
+ ); + }, + size: 120, + }, + { + header: "Mode", + accessorKey: "mode", + enableSorting: true, + cell: ({ row }) => { + const mode = row.original.mode; + const getModeIcon = (mode: string) => { + switch (mode?.toLowerCase()) { + case "chat": + return "💬"; + case "rerank": + return "🔄"; + case "embedding": + return "📄"; + default: + return "🤖"; + } + }; + + return ( +
+ {getModeIcon(mode || "")} + {mode || "Chat"} +
+ ); + }, + size: 100, + }, + { + header: "Max Input", + accessorKey: "max_input_tokens", + enableSorting: true, + cell: ({ row }) => ( + {formatTokens(row.original.max_input_tokens)} + ), + size: 100, + meta: { + className: "text-center", + }, + }, + { + header: "Max Output", + accessorKey: "max_output_tokens", + enableSorting: true, + cell: ({ row }) => ( + {formatTokens(row.original.max_output_tokens)} + ), + size: 100, + meta: { + className: "text-center", + }, + }, + { + header: "Input $/1M", + accessorKey: "input_cost_per_token", + enableSorting: true, + cell: ({ row }) => { + const cost = row.original.input_cost_per_token; + return ( + + {cost ? formatCost(cost) : "Free"} + + ); + }, + size: 100, + meta: { + className: "text-center", + }, + }, + { + header: "Output $/1M", + accessorKey: "output_cost_per_token", + enableSorting: true, + cell: ({ row }) => { + const cost = row.original.output_cost_per_token; + return ( + + {cost ? formatCost(cost) : "Free"} + + ); + }, + size: 100, + meta: { + className: "text-center", + }, + }, + { + header: "Features", + accessorKey: "supports_vision", + enableSorting: false, + cell: ({ row }) => { + const model = row.original; + + // Dynamically get all features that start with 'supports_' and are true + const features = Object.entries(model) + .filter(([key, value]) => key.startsWith('supports_') && value === true) + .map(([key]) => formatCapabilityName(key)); + + if (features.length === 0) { + return -; + } + + if (features.length === 1) { + return ( +
+ + {features[0]} + +
+ ); + } + + return ( +
+ + {features[0]} + + +
All Features:
+ {features.map((feature, index) => ( +
• {feature}
+ ))} +
+ } + trigger="click" + placement="topLeft" + > + e.stopPropagation()} + > + +{features.length - 1} + + +
+ ); + }, + size: 120, + }, + { + header: "Limits", + accessorKey: "rpm", + enableSorting: true, + cell: ({ row }) => { + const model = row.original; + return ( + + {formatLimits(model.rpm, model.tpm)} + + ); + }, + size: 150, + }, + ]; + + return ( +
+ {/* Header */} +
+
+ {pageTitle} +
+
+ +
+ {/* About Section */} + + About +

{customDocsDescription ? customDocsDescription : "Proxy Server to call 100+ LLMs in the OpenAI format."}

+
+ + 🔧 + Built with litellm: v{litellmVersion} + +
+
+ + {/* Useful Links */} + {usefulLinks && Object.keys(usefulLinks).length > 0 && ( + + Useful Links +
+ {Object.entries(usefulLinks || {}).map(([title, url]) => ( + + ))} +
+
+ )} + + {/* Health and Endpoint Status */} + + Health and Endpoint Status +
+ Service status: {serviceStatus} +
+
+ + {/* Filters */} + +
+
+
+ Search Models: + + + +
+
+ + setSearchTerm(e.target.value)} + className="border rounded-lg pl-10 pr-4 py-3 w-full text-base focus:outline-none focus:ring-2 focus:ring-blue-500 focus:border-transparent" + /> +
+
+
+ Provider: + +
+
+ Mode: + +
+
+ Features: + +
+
+
+ + {/* Models Table */} + +
+ Available Models +
+ + + +
+ + Showing {filteredData.length} of {modelHubData?.length || 0} models + +
+
+
+ + {/* Model Details Modal */} + + {selectedModel?.model_group || "Model Details"} + {selectedModel && ( + + copyToClipboard(selectedModel.model_group)} + className="cursor-pointer text-gray-500 hover:text-blue-500 w-4 h-4" + /> + + )} +
+ } + width={1000} + open={isModalVisible} + footer={null} + onOk={handleModalOk} + onCancel={handleModalCancel} + > + {selectedModel && ( +
+ {/* Model Overview */} +
+ Model Overview +
+
+ Model Name: + {selectedModel.model_group} +
+
+ Mode: + {selectedModel.mode || "Not specified"} +
+
+ Providers: +
+ {selectedModel.providers.map(provider => { + const { logo } = getProviderLogoAndName(provider); + return ( + +
+ {logo && ( + {provider} { + (e.target as HTMLImageElement).style.display = 'none'; + }} + /> + )} + {provider} +
+
+ ); + })} +
+
+
+ + {/* Wildcard Routing Note */} + {selectedModel.model_group.includes('*') && ( +
+
+ +
+ Wildcard Routing + + This model uses wildcard routing. You can pass any value where you see the * symbol. + + + For example, with {selectedModel.model_group}, you can use any string ({selectedModel.model_group.replace('*', 'my-custom-value')}) that matches this pattern. + +
+
+
+ )} +
+ + {/* Token and Cost Information */} +
+ Token & Cost Information +
+
+ Max Input Tokens: + {selectedModel.max_input_tokens?.toLocaleString() || "Not specified"} +
+
+ Max Output Tokens: + {selectedModel.max_output_tokens?.toLocaleString() || "Not specified"} +
+
+ Input Cost per 1M Tokens: + {selectedModel.input_cost_per_token ? formatCost(selectedModel.input_cost_per_token) : "Not specified"} +
+
+ Output Cost per 1M Tokens: + {selectedModel.output_cost_per_token ? formatCost(selectedModel.output_cost_per_token) : "Not specified"} +
+
+
+ + {/* Capabilities */} +
+ Capabilities +
+ {(() => { + const capabilities = getModelCapabilities(selectedModel); + const colors = ['green', 'blue', 'purple', 'orange', 'red', 'yellow']; + + if (capabilities.length === 0) { + return No special capabilities listed; + } + + return capabilities.map((capability, index) => ( + + {formatCapabilityName(capability)} + + )); + })()} +
+
+ + {/* Rate Limits */} + {(selectedModel.tpm || selectedModel.rpm) && ( +
+ Rate Limits +
+ {selectedModel.tpm && ( +
+ Tokens per Minute: + {selectedModel.tpm.toLocaleString()} +
+ )} + {selectedModel.rpm && ( +
+ Requests per Minute: + {selectedModel.rpm.toLocaleString()} +
+ )} +
+
+ )} + + {/* Supported OpenAI Parameters */} + {selectedModel.supported_openai_params && ( +
+ Supported OpenAI Parameters +
+ {selectedModel.supported_openai_params.map(param => ( + {param} + ))} +
+
+ )} + + {/* Usage Example */} +
+ Usage Example +
+
+{(() => {
+  const codeSnippet = generateCodeSnippet({
+    apiKeySource: 'custom',
+    accessToken: null,
+    apiKey: 'your_api_key',
+    inputMessage: 'Hello, how are you?',
+    chatHistory: [
+      { role: 'user', content: 'Hello, how are you?', isImage: false } as MessageType
+    ],
+    selectedTags: [],
+    selectedVectorStores: [],
+    selectedGuardrails: [],
+    endpointType: getEndpointType(selectedModel.mode || 'chat'),
+    selectedModel: selectedModel.model_group,
+    selectedSdk: 'openai'
+  });
+  return codeSnippet;
+})()}
+                
+
+
+ +
+
+
+ )} + + + ); +}; + +export default PublicModelHub; \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/public_model_hub_columns.tsx b/ui/litellm-dashboard/src/components/public_model_hub_columns.tsx new file mode 100644 index 00000000000..faa965e60a2 --- /dev/null +++ b/ui/litellm-dashboard/src/components/public_model_hub_columns.tsx @@ -0,0 +1,218 @@ +import React from "react"; +import { ColumnDef } from "@tanstack/react-table"; +import { Badge, Text } from "@tremor/react"; +import { EyeIcon, CogIcon } from "@heroicons/react/outline"; +import { Tag } from "antd"; + +interface ModelGroupInfo { + model_group: string; + providers: string[]; + max_input_tokens?: number; + max_output_tokens?: number; + input_cost_per_token?: number; + output_cost_per_token?: number; + mode?: string; + tpm?: number; + rpm?: number; + supports_parallel_function_calling: boolean; + supports_vision: boolean; + supports_function_calling: boolean; + supported_openai_params?: string[]; + [key: string]: any; +} + +const formatCost = (cost: number) => { + return `$${(cost * 1_000_000).toFixed(4)}`; +}; + +const formatTokens = (tokens: number | undefined) => { + if (!tokens) return "N/A"; + if (tokens >= 1000) { + return `${(tokens / 1000).toFixed(0)}K`; + } + return tokens.toString(); +}; + +const formatLimits = (rpm?: number, tpm?: number) => { + const limits = []; + if (rpm) limits.push(`RPM: ${rpm.toLocaleString()}`); + if (tpm) limits.push(`TPM: ${tpm.toLocaleString()}`); + return limits.length > 0 ? limits.join(", ") : "N/A"; +}; + +export const publicModelHubColumns = (): ColumnDef[] => [ + { + header: "#", + id: "index", + enableSorting: false, + cell: ({ row }) => { + const index = row.index + 1; + return {index}; + }, + size: 50, + }, + { + header: "Model Group", + accessorKey: "model_group", + enableSorting: true, + cell: ({ row }) => ( + {row.original.model_group} + ), + size: 150, + }, + { + header: "Providers", + accessorKey: "providers", + enableSorting: true, + cell: ({ row }) => { + const providers = row.original.providers; + const getProviderColor = (provider: string) => { + switch (provider.toLowerCase()) { + case "openai": + return "green"; + case "anthropic": + return "orange"; + case "cohere": + return "blue"; + default: + return "gray"; + } + }; + + return ( +
+ {providers.map((provider) => ( + + {provider} + + ))} +
+ ); + }, + size: 120, + }, + { + header: "Mode", + accessorKey: "mode", + enableSorting: true, + cell: ({ row }) => { + const mode = row.original.mode; + const getModeIcon = (mode: string) => { + switch (mode?.toLowerCase()) { + case "chat": + return "💬"; + case "rerank": + return "🔄"; + case "embedding": + return "📄"; + default: + return "🤖"; + } + }; + + return ( +
+ {getModeIcon(mode || "")} + {mode || "Chat"} +
+ ); + }, + size: 100, + }, + { + header: "Max Input", + accessorKey: "max_input_tokens", + enableSorting: true, + cell: ({ row }) => ( + {formatTokens(row.original.max_input_tokens)} + ), + size: 100, + }, + { + header: "Max Output", + accessorKey: "max_output_tokens", + enableSorting: true, + cell: ({ row }) => ( + {formatTokens(row.original.max_output_tokens)} + ), + size: 100, + }, + { + header: "Input $/1K", + accessorKey: "input_cost_per_token", + enableSorting: true, + cell: ({ row }) => { + const cost = row.original.input_cost_per_token; + return ( + + {cost ? formatCost(cost) : "Free"} + + ); + }, + size: 100, + }, + { + header: "Output $/1K", + accessorKey: "output_cost_per_token", + enableSorting: true, + cell: ({ row }) => { + const cost = row.original.output_cost_per_token; + return ( + + {cost ? formatCost(cost) : "Free"} + + ); + }, + size: 100, + }, + { + header: "Features", + accessorKey: "supports_vision", + enableSorting: false, + cell: ({ row }) => { + const model = row.original; + const features = []; + + if (model.supports_vision) { + features.push( +
+ +
+ ); + } + + if (model.supports_function_calling || model.supports_parallel_function_calling) { + features.push( +
+ +
+ ); + } + + return features.length > 0 ? ( +
{features}
+ ) : ( + - + ); + }, + size: 100, + }, + { + header: "Limits", + accessorKey: "rpm", + enableSorting: true, + cell: ({ row }) => { + const model = row.original; + return ( + + {formatLimits(model.rpm, model.tpm)} + + ); + }, + size: 150, + }, +]; \ No newline at end of file