diff --git a/litellm/main.py b/litellm/main.py index 1840a900a60..6200d8cb18a 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -3387,7 +3387,7 @@ def embedding( client=client, aembedding=aembedding, ) - elif custom_llm_provider == "cohere": + elif custom_llm_provider == "cohere" or custom_llm_provider == "cohere_chat": cohere_key = ( api_key or litellm.cohere_key @@ -3404,7 +3404,7 @@ def embedding( logging_obj=logging, model_response=EmbeddingResponse(), aembedding=aembedding, - timeout=float(timeout), + timeout=timeout, client=client, ) elif custom_llm_provider == "huggingface": diff --git a/litellm/tests/test_embedding.py b/litellm/tests/test_embedding.py index 9c1e4aa2cd6..31268395f14 100644 --- a/litellm/tests/test_embedding.py +++ b/litellm/tests/test_embedding.py @@ -282,19 +282,19 @@ async def test_cohere_embedding(sync_mode): # test_cohere_embedding() -def test_cohere_embedding3(): +@pytest.mark.parametrize("custom_llm_provider", ["cohere", "cohere_chat"]) +@pytest.mark.asyncio() +async def test_cohere_embedding3(custom_llm_provider): try: litellm.set_verbose = True - response = embedding( - model="embed-english-v3.0", + response = await litellm.aembedding( + model=f"{custom_llm_provider}/embed-english-v3.0", input=["good morning from litellm", "this is another item"], + timeout=None, + max_retries=0, ) print(f"response:", response) - custom_llm_provider = response._hidden_params["custom_llm_provider"] - - assert custom_llm_provider == "cohere" - except Exception as e: pytest.fail(f"Error occurred: {e}") diff --git a/ui/litellm-dashboard/src/components/model_dashboard.tsx b/ui/litellm-dashboard/src/components/model_dashboard.tsx index 2d4fbe36f5b..de92805b0bf 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard.tsx @@ -147,6 +147,7 @@ enum Providers { MistralAI = "Mistral AI", OpenAI_Compatible = "OpenAI-Compatible Endpoints (Together AI, etc.)", Vertex_AI = "Vertex AI (Anthropic, Gemini, etc.)", + Cohere = "Cohere", Databricks = "Databricks", Ollama = "Ollama", } @@ -160,6 +161,7 @@ const provider_map: Record = { Bedrock: "bedrock", Groq: "groq", MistralAI: "mistral", + Cohere: "cohere_chat", OpenAI_Compatible: "openai", Vertex_AI: "vertex_ai", Databricks: "databricks", @@ -928,7 +930,26 @@ const ModelDashboard: React.FC = ({ _providerModels.push(key); } }); + + // Special case for cohere_chat + // we need both cohere_chat and cohere models to show on dropdown + if (providerKey == Providers.Cohere) { + console.log("adding cohere chat model") + Object.entries(modelMap).forEach(([key, value]) => { + if ( + value !== null && + typeof value === "object" && + "litellm_provider" in (value as object) && + ((value as any)["litellm_provider"] === "cohere") + ) { + _providerModels.push(key); + } + }); + } } + + + setProviderModels(_providerModels); console.log(`providerModels: ${providerModels}`); } @@ -1785,7 +1806,7 @@ const ModelDashboard: React.FC = ({