Skip to content
Draft
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
26 changes: 23 additions & 3 deletions flow/flow/doctype/flow_model/flow_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,21 +139,41 @@ def test_connection(self):
kwargs = {
"model": self.model_id,
"api_key": api_key,
"messages": [{"role": "user", "content": "ping"}],
"max_tokens": 1,
"timeout": 15,
}
if base_url:
kwargs["api_base"] = base_url

try:
litellm.completion(**kwargs)
if _is_embedding_model(self.model_id):
litellm.embedding(input=["ping"], encoding_format="float", **kwargs)
else:
litellm.completion(
messages=[{"role": "user", "content": "ping"}],
max_tokens=1,
**kwargs,
)
except Exception as e:
frappe.throw(str(e)[:500] or type(e).__name__, title=_(type(e).__name__))

return {"ok": True, "message": _("Connection OK")}


def _is_embedding_model(model_id: str) -> bool:
"""Whether ``model_id`` should be tested through the embeddings endpoint."""
import litellm

try:
if litellm.get_model_info(model_id).get("mode") == "embedding":
return True
except Exception:
# Custom OpenAI-compatible models are often absent from LiteLLM's registry.
pass

model_name = model_id.rsplit("/", 1)[-1].lower()
return "embedding" in model_name


def _detect_context_window(model_id: str) -> int:
"""Max input tokens for `model_id` per litellm, or 0 if unknown/unmapped."""
import litellm
Expand Down
63 changes: 62 additions & 1 deletion flow/flow/doctype/flow_model/test_flow_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import frappe
from frappe.tests import IntegrationTestCase

from flow.flow.doctype.flow_model.flow_model import _detect_context_window
from flow.flow.doctype.flow_model.flow_model import FlowModel, _detect_context_window, _is_embedding_model


def _model(**overrides: Any) -> dict:
Expand Down Expand Up @@ -180,3 +180,64 @@ def test_keeps_existing_when_new_model_unmapped(self):
doc.model_id = "anthropic/made-up-model-xyz"
doc.save()
self.assertEqual(doc.context_window, 128000)


class TestFlowModelConnection(IntegrationTestCase):
def tearDown(self):
frappe.db.rollback()

def test_detects_mapped_embedding_model(self):
with patch("litellm.get_model_info", return_value={"mode": "embedding"}):
self.assertTrue(_is_embedding_model("vendor/vector-model"))

def test_detects_unmapped_embedding_model_by_name(self):
with patch("litellm.get_model_info", side_effect=Exception("not mapped")):
self.assertTrue(_is_embedding_model("openai/text-embedding-v4"))

def test_connection_uses_embedding_endpoint_for_unmapped_embedding_model(self):
doc = frappe.get_doc(
_model(model_id="openai/text-embedding-v4", base_url="https://api.example.com/v1")
)

with (
patch.object(FlowModel, "check_permission"),
patch.object(FlowModel, "get_password", return_value="sk-test"),
patch("flow.lib.model.resolve_provider_credentials", return_value={}),
patch("litellm.get_model_info", side_effect=Exception("not mapped")),
patch("litellm.embedding") as embedding,
patch("litellm.completion") as completion,
):
result = doc.test_connection()

embedding.assert_called_once_with(
model="openai/text-embedding-v4",
api_key="sk-test",
timeout=15,
api_base="https://api.example.com/v1",
input=["ping"],
encoding_format="float",
)
completion.assert_not_called()
self.assertTrue(result["ok"])

def test_connection_keeps_chat_models_on_completion_endpoint(self):
doc = frappe.get_doc(_model(model_id="openai/gpt-4o-mini"))

with (
patch.object(FlowModel, "check_permission"),
patch.object(FlowModel, "get_password", return_value="sk-test"),
patch("flow.lib.model.resolve_provider_credentials", return_value={}),
patch("litellm.get_model_info", return_value={"mode": "chat"}),
patch("litellm.embedding") as embedding,
patch("litellm.completion") as completion,
):
doc.test_connection()

embedding.assert_not_called()
completion.assert_called_once_with(
model="openai/gpt-4o-mini",
api_key="sk-test",
timeout=15,
messages=[{"role": "user", "content": "ping"}],
max_tokens=1,
)