mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
chore(api): Upgrade graphon to v0.5.1 (#37168)
Co-authored-by: Yunlu Wen <yunlu.wen@dify.ai> Co-authored-by: L1nSn0w <l1nsn0w@qq.com>
This commit is contained in:
co-authored by
Yunlu Wen
L1nSn0w
parent
fb39df49c8
commit
b4c50eb920
@@ -0,0 +1,138 @@
|
||||
app:
|
||||
description: Response stream ordering fixture matching graphon issue 170.
|
||||
icon: 🤖
|
||||
icon_background: '#FFEAD5'
|
||||
mode: advanced-chat
|
||||
name: response_stream_filter_issue_170_workflow
|
||||
use_icon_as_answer_icon: false
|
||||
dependencies: []
|
||||
kind: app
|
||||
version: 0.3.1
|
||||
workflow:
|
||||
conversation_variables: []
|
||||
environment_variables: []
|
||||
features:
|
||||
file_upload: {}
|
||||
opening_statement: ''
|
||||
retriever_resource:
|
||||
enabled: true
|
||||
sensitive_word_avoidance:
|
||||
enabled: false
|
||||
speech_to_text:
|
||||
enabled: false
|
||||
suggested_questions: []
|
||||
suggested_questions_after_answer:
|
||||
enabled: false
|
||||
text_to_speech:
|
||||
enabled: false
|
||||
language: ''
|
||||
voice: ''
|
||||
graph:
|
||||
edges:
|
||||
- id: start-llm
|
||||
source: start
|
||||
sourceHandle: source
|
||||
target: llm
|
||||
targetHandle: target
|
||||
- id: llm-dufu
|
||||
source: llm
|
||||
sourceHandle: source
|
||||
target: dufu
|
||||
targetHandle: target
|
||||
- id: dufu-answer
|
||||
source: dufu
|
||||
sourceHandle: source
|
||||
target: answer
|
||||
targetHandle: target
|
||||
nodes:
|
||||
- data:
|
||||
desc: ''
|
||||
title: Start
|
||||
type: start
|
||||
variables: []
|
||||
id: start
|
||||
position:
|
||||
x: 80
|
||||
y: 282
|
||||
sourcePosition: right
|
||||
targetPosition: left
|
||||
type: custom
|
||||
- data:
|
||||
context:
|
||||
enabled: false
|
||||
variable_selector: []
|
||||
desc: ''
|
||||
memory:
|
||||
query_prompt_template: '{{#sys.query#}}'
|
||||
window:
|
||||
enabled: false
|
||||
size: 10
|
||||
model:
|
||||
completion_params:
|
||||
temperature: 0.7
|
||||
mode: chat
|
||||
name: gpt-4o-mini
|
||||
provider: openai
|
||||
prompt_template:
|
||||
- role: system
|
||||
text: Please output a poem by Li Bai
|
||||
selected: false
|
||||
title: Li Bai
|
||||
type: llm
|
||||
variables: []
|
||||
vision:
|
||||
enabled: false
|
||||
id: llm
|
||||
position:
|
||||
x: 380
|
||||
y: 282
|
||||
sourcePosition: right
|
||||
targetPosition: left
|
||||
type: custom
|
||||
- data:
|
||||
context:
|
||||
enabled: false
|
||||
variable_selector: []
|
||||
desc: ''
|
||||
model:
|
||||
completion_params:
|
||||
temperature: 0.7
|
||||
mode: chat
|
||||
name: gpt-4o-mini
|
||||
provider: openai
|
||||
prompt_template:
|
||||
- role: system
|
||||
text: Please output a poem by Du Fu
|
||||
selected: false
|
||||
title: Du Fu
|
||||
type: llm
|
||||
variables: []
|
||||
vision:
|
||||
enabled: false
|
||||
id: dufu
|
||||
position:
|
||||
x: 680
|
||||
y: 282
|
||||
sourcePosition: right
|
||||
targetPosition: left
|
||||
type: custom
|
||||
- data:
|
||||
answer: |-
|
||||
# Du Fu
|
||||
|
||||
{{#dufu.text#}}
|
||||
|
||||
# Li Bai
|
||||
|
||||
{{#llm.text#}}
|
||||
desc: ''
|
||||
title: Answer
|
||||
type: answer
|
||||
variables: []
|
||||
id: answer
|
||||
position:
|
||||
x: 980
|
||||
y: 282
|
||||
sourcePosition: right
|
||||
targetPosition: left
|
||||
type: custom
|
||||
@@ -0,0 +1,75 @@
|
||||
"""Integration coverage for Dify's ResponseStreamFilter boundary behavior."""
|
||||
|
||||
from core.workflow.workflow_entry import iter_dify_graph_engine_events
|
||||
from graphon.graph_engine import GraphEngine, GraphEngineConfig
|
||||
from graphon.graph_engine.command_channels import InMemoryChannel
|
||||
from graphon.graph_events import GraphRunSucceededEvent, NodeRunStreamChunkEvent
|
||||
from tests.unit_tests.core.workflow.graph_engine.test_mock_config import MockConfigBuilder
|
||||
from tests.unit_tests.core.workflow.graph_engine.test_table_runner import WorkflowRunner
|
||||
|
||||
|
||||
def _build_issue_170_mock_config():
|
||||
runner = WorkflowRunner()
|
||||
mock_config = (
|
||||
MockConfigBuilder()
|
||||
.with_node_output(
|
||||
"llm",
|
||||
{
|
||||
"text": "Quiet Night Thought",
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
)
|
||||
.with_node_output(
|
||||
"dufu",
|
||||
{
|
||||
"text": "Spring View",
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
)
|
||||
.build()
|
||||
)
|
||||
|
||||
return runner, mock_config
|
||||
|
||||
|
||||
def test_dify_response_stream_filter_handles_issue_170_shape() -> None:
|
||||
runner, mock_config = _build_issue_170_mock_config()
|
||||
fixture_data = runner.load_fixture("response_stream_filter_issue_170_workflow")
|
||||
graph, graph_runtime_state = runner.create_graph_from_fixture(
|
||||
fixture_data=fixture_data,
|
||||
query="1",
|
||||
use_mock_factory=True,
|
||||
mock_config=mock_config,
|
||||
)
|
||||
|
||||
expected_answer = "# Du Fu\n\nSpring View\n\n# Li Bai\n\nQuiet Night Thought"
|
||||
|
||||
engine = GraphEngine(
|
||||
workflow_id="test_workflow",
|
||||
graph=graph,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
command_channel=InMemoryChannel(),
|
||||
config=GraphEngineConfig(),
|
||||
)
|
||||
events = list(iter_dify_graph_engine_events(engine))
|
||||
|
||||
stream_chunk_events = [event for event in events if isinstance(event, NodeRunStreamChunkEvent)]
|
||||
success_events = [event for event in events if isinstance(event, GraphRunSucceededEvent)]
|
||||
|
||||
assert success_events
|
||||
assert stream_chunk_events
|
||||
actual_answer = "".join(event.chunk for event in stream_chunk_events)
|
||||
assert actual_answer.strip() == expected_answer
|
||||
assert stream_chunk_events[-1].is_final is True
|
||||
assert success_events[-1].outputs["answer"].strip() == expected_answer
|
||||
assert actual_answer.strip() == success_events[-1].outputs["answer"].strip()
|
||||
+9
-5
@@ -9,7 +9,7 @@ from core.repositories.human_input_repository import (
|
||||
HumanInputFormEntity,
|
||||
HumanInputFormRepository,
|
||||
)
|
||||
from core.workflow.node_runtime import DifyFileReferenceFactory, DifyHumanInputNodeRuntime
|
||||
from core.workflow.node_runtime import DifyHumanInputNodeRuntime
|
||||
from core.workflow.system_variables import build_system_variables
|
||||
from graphon.entities import WorkflowStartReason
|
||||
from graphon.file import File, FileTransferMethod, FileType
|
||||
@@ -186,25 +186,29 @@ def _build_graph(runtime_state: GraphRuntimeState, repo: HumanInputFormRepositor
|
||||
)
|
||||
|
||||
human_a_config = {"id": "human_a", "data": human_data.model_dump()}
|
||||
human_a_runtime = DifyHumanInputNodeRuntime(graph_init_params.run_context)
|
||||
human_a_runtime._file_reference_factory = _TestFileReferenceFactory() # type: ignore[attr-defined]
|
||||
human_a = HumanInputNode(
|
||||
node_id=human_a_config["id"],
|
||||
data=human_data,
|
||||
graph_init_params=graph_init_params,
|
||||
graph_runtime_state=runtime_state,
|
||||
form_repository=repo,
|
||||
file_reference_factory=DifyFileReferenceFactory(graph_init_params.run_context),
|
||||
runtime=DifyHumanInputNodeRuntime(graph_init_params.run_context),
|
||||
file_reference_factory=_TestFileReferenceFactory(),
|
||||
runtime=human_a_runtime,
|
||||
)
|
||||
|
||||
human_b_config = {"id": "human_b", "data": human_data.model_dump()}
|
||||
human_b_runtime = DifyHumanInputNodeRuntime(graph_init_params.run_context)
|
||||
human_b_runtime._file_reference_factory = _TestFileReferenceFactory() # type: ignore[attr-defined]
|
||||
human_b = HumanInputNode(
|
||||
node_id=human_b_config["id"],
|
||||
data=human_data,
|
||||
graph_init_params=graph_init_params,
|
||||
graph_runtime_state=runtime_state,
|
||||
form_repository=repo,
|
||||
file_reference_factory=DifyFileReferenceFactory(graph_init_params.run_context),
|
||||
runtime=DifyHumanInputNodeRuntime(graph_init_params.run_context),
|
||||
file_reference_factory=_TestFileReferenceFactory(),
|
||||
runtime=human_b_runtime,
|
||||
)
|
||||
|
||||
end_data = EndNodeData(
|
||||
|
||||
@@ -24,6 +24,7 @@ from core.tools.utils.yaml_utils import _load_yaml_file
|
||||
from core.workflow.node_factory import DifyNodeFactory, get_default_root_node_id
|
||||
from core.workflow.system_variables import build_bootstrap_variables, build_system_variables
|
||||
from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool
|
||||
from core.workflow.workflow_entry import iter_dify_graph_engine_events
|
||||
from graphon.entities import GraphInitParams
|
||||
from graphon.graph import Graph
|
||||
from graphon.graph_engine import GraphEngine, GraphEngineConfig
|
||||
@@ -386,7 +387,7 @@ class TableTestRunner:
|
||||
|
||||
# Execute and collect events
|
||||
events: list[GraphEngineEvent] = []
|
||||
for event in engine.run():
|
||||
for event in iter_dify_graph_engine_events(engine):
|
||||
events.append(event)
|
||||
|
||||
# Check execution success
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from core.workflow.workflow_entry import iter_dify_graph_engine_events
|
||||
from graphon.graph_engine import GraphEngine, GraphEngineConfig
|
||||
from graphon.graph_engine.command_channels import InMemoryChannel
|
||||
from graphon.graph_events import (
|
||||
@@ -31,20 +32,17 @@ def test_tool_in_chatflow():
|
||||
config=GraphEngineConfig(),
|
||||
)
|
||||
|
||||
events = list(engine.run())
|
||||
events = list(iter_dify_graph_engine_events(engine))
|
||||
|
||||
# Check for successful completion
|
||||
success_events = [e for e in events if isinstance(e, GraphRunSucceededEvent)]
|
||||
assert len(success_events) > 0, "Workflow should complete successfully"
|
||||
|
||||
# Check for streaming events
|
||||
stream_chunk_events = [e for e in events if isinstance(e, NodeRunStreamChunkEvent)]
|
||||
stream_chunk_count = len(stream_chunk_events)
|
||||
|
||||
assert stream_chunk_count == 1, f"Expected 1 streaming events, but got {stream_chunk_count}"
|
||||
assert stream_chunk_events[0].chunk == "hello, dify!", (
|
||||
f"Expected chunk to be 'hello, dify!', but got {stream_chunk_events[0].chunk}"
|
||||
)
|
||||
assert len(stream_chunk_events) > 0
|
||||
assert "".join(event.chunk for event in stream_chunk_events) == "hello, dify!"
|
||||
assert stream_chunk_events[-1].is_final is True
|
||||
assert success_events[-1].outputs["answer"] == "hello, dify!"
|
||||
|
||||
|
||||
def test_answer_can_render_llm_structured_output_in_chatflow():
|
||||
@@ -88,7 +86,7 @@ def test_answer_can_render_llm_structured_output_in_chatflow():
|
||||
config=GraphEngineConfig(),
|
||||
)
|
||||
|
||||
events = list(engine.run())
|
||||
events = list(iter_dify_graph_engine_events(engine))
|
||||
success_events = [e for e in events if isinstance(e, GraphRunSucceededEvent)]
|
||||
|
||||
assert success_events, "Workflow should complete successfully"
|
||||
|
||||
@@ -30,7 +30,7 @@ from core.workflow.human_input_adapter import (
|
||||
WebAppDeliveryMethod,
|
||||
_WebAppDeliveryConfig,
|
||||
)
|
||||
from core.workflow.node_runtime import DifyFileReferenceFactory, DifyHumanInputNodeRuntime
|
||||
from core.workflow.node_runtime import DifyHumanInputNodeRuntime
|
||||
from core.workflow.system_variables import build_system_variables
|
||||
from graphon.entities import GraphInitParams
|
||||
from graphon.file import File, FileTransferMethod, FileType
|
||||
@@ -171,12 +171,13 @@ def _build_human_input_node(
|
||||
typed_node_data = (
|
||||
node_data if isinstance(node_data, HumanInputNodeData) else HumanInputNodeData.model_validate(node_data)
|
||||
)
|
||||
runtime._file_reference_factory = _TestFileReferenceFactory() # type: ignore[attr-defined]
|
||||
return HumanInputNode(
|
||||
node_id=node_id,
|
||||
data=typed_node_data,
|
||||
graph_init_params=graph_init_params,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
file_reference_factory=DifyFileReferenceFactory(graph_init_params.run_context),
|
||||
file_reference_factory=_TestFileReferenceFactory(),
|
||||
runtime=runtime,
|
||||
)
|
||||
|
||||
|
||||
+5
-3
@@ -4,7 +4,7 @@ from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, InvokeFrom, UserFrom
|
||||
from core.workflow.node_runtime import DifyFileReferenceFactory, DifyHumanInputNodeRuntime
|
||||
from core.workflow.node_runtime import DifyHumanInputNodeRuntime
|
||||
from core.workflow.system_variables import default_system_variables
|
||||
from graphon.entities import GraphInitParams
|
||||
from graphon.enums import BuiltinNodeTypes
|
||||
@@ -67,14 +67,16 @@ def _create_human_input_node(
|
||||
if isinstance(config["data"], HumanInputNodeData)
|
||||
else HumanInputNodeData.model_validate(config["data"])
|
||||
)
|
||||
runtime = DifyHumanInputNodeRuntime(graph_init_params.run_context)
|
||||
runtime._file_reference_factory = _TestFileReferenceFactory() # type: ignore[attr-defined]
|
||||
return HumanInputNode(
|
||||
node_id=config["id"],
|
||||
data=node_data,
|
||||
graph_init_params=graph_init_params,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
form_repository=repo,
|
||||
file_reference_factory=DifyFileReferenceFactory(graph_init_params.run_context),
|
||||
runtime=DifyHumanInputNodeRuntime(graph_init_params.run_context),
|
||||
file_reference_factory=_TestFileReferenceFactory(),
|
||||
runtime=runtime,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -76,6 +76,7 @@ from graphon.nodes.llm.node import (
|
||||
_render_jinja2_message,
|
||||
)
|
||||
from graphon.nodes.llm.protocols import CredentialsProvider, ModelFactory
|
||||
from graphon.nodes.llm.reasoning import split_reasoning
|
||||
from graphon.nodes.llm.runtime_protocols import PromptMessageSerializerProtocol
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
from graphon.template_rendering import TemplateRenderError
|
||||
@@ -1271,7 +1272,10 @@ class TestLLMNodeSaveMultiModalImageOutput:
|
||||
assert llm_node._file_outputs == [mock_file]
|
||||
assert file == mock_file
|
||||
mock_file_saver.save_binary_string.assert_called_once_with(
|
||||
data=b"test-data", mime_type="image/png", file_type=FileType.IMAGE
|
||||
data=b"test-data",
|
||||
mime_type="image/png",
|
||||
file_type=FileType.IMAGE,
|
||||
extension_override=".png",
|
||||
)
|
||||
|
||||
def test_llm_node_save_url_output(self, llm_node_for_multimodal: tuple[LLMNode, LLMFileSaver]):
|
||||
@@ -1305,8 +1309,9 @@ class TestLLMNodeSaveMultiModalImageOutput:
|
||||
|
||||
def test_llm_node_image_file_to_markdown(llm_node: LLMNode):
|
||||
mock_file = mock.MagicMock(spec=File)
|
||||
mock_file.type = FileType.IMAGE
|
||||
mock_file.generate_url.return_value = "https://example.com/image.png"
|
||||
markdown = llm_node._image_file_to_markdown(mock_file)
|
||||
markdown = llm_node._saved_file_to_markdown(mock_file)
|
||||
assert markdown == ""
|
||||
|
||||
|
||||
@@ -1378,6 +1383,7 @@ class TestSaveMultimodalOutputAndConvertResultToMarkdown:
|
||||
data=image_raw_data,
|
||||
mime_type="image/png",
|
||||
file_type=FileType.IMAGE,
|
||||
extension_override=".png",
|
||||
)
|
||||
assert mock_saved_file in llm_node._file_outputs
|
||||
|
||||
@@ -1425,7 +1431,7 @@ class TestReasoningFormat:
|
||||
</think>Dify is an open source AI platform.
|
||||
"""
|
||||
|
||||
clean_text, reasoning_content = LLMNode._split_reasoning(text_with_think, "separated")
|
||||
clean_text, reasoning_content = split_reasoning(text_with_think, "separated")
|
||||
|
||||
assert clean_text == "Dify is an open source AI platform."
|
||||
assert reasoning_content == "I need to explain what Dify is. It's an open source AI platform."
|
||||
@@ -1438,7 +1444,7 @@ class TestReasoningFormat:
|
||||
</think>Dify is an open source AI platform.
|
||||
"""
|
||||
|
||||
clean_text, reasoning_content = LLMNode._split_reasoning(text_with_think, "tagged")
|
||||
clean_text, reasoning_content = split_reasoning(text_with_think, "tagged")
|
||||
|
||||
# Original text unchanged
|
||||
assert clean_text == text_with_think
|
||||
@@ -1450,7 +1456,7 @@ class TestReasoningFormat:
|
||||
|
||||
text_without_think = "This is a simple answer without any thinking blocks."
|
||||
|
||||
clean_text, reasoning_content = LLMNode._split_reasoning(text_without_think, "separated")
|
||||
clean_text, reasoning_content = split_reasoning(text_without_think, "separated")
|
||||
|
||||
assert clean_text == text_without_think
|
||||
assert reasoning_content == ""
|
||||
@@ -1471,7 +1477,7 @@ class TestReasoningFormat:
|
||||
<think>I need to explain what Dify is. It's an open source AI platform.
|
||||
</think>Dify is an open source AI platform.
|
||||
"""
|
||||
clean_text, reasoning_content = LLMNode._split_reasoning(text_with_think, node_data.reasoning_format)
|
||||
clean_text, reasoning_content = split_reasoning(text_with_think, node_data.reasoning_format)
|
||||
|
||||
assert clean_text == text_with_think
|
||||
assert reasoning_content == ""
|
||||
@@ -1569,10 +1575,10 @@ def test_handle_invoke_result_streaming_collects_text_metrics_and_structured_out
|
||||
)
|
||||
|
||||
assert events[0] == first_chunk
|
||||
assert events[1] == StreamChunkEvent(selector=["node-1", "text"], chunk="<think>plan</think>", is_final=False)
|
||||
assert events[2] == StreamChunkEvent(selector=["node-1", "text"], chunk="answer", is_final=False)
|
||||
|
||||
completed = events[3]
|
||||
assert events[1] == StreamChunkEvent(selector=["node-1", "text"], chunk="answer", is_final=False)
|
||||
|
||||
completed = events[2]
|
||||
assert isinstance(completed, ModelInvokeCompletedEvent)
|
||||
assert completed.text == "answer"
|
||||
assert completed.reasoning_content == "plan"
|
||||
|
||||
@@ -338,6 +338,52 @@ class TestWorkflowEntryRun:
|
||||
|
||||
assert list(entry.run()) == []
|
||||
|
||||
def test_iter_dify_graph_engine_events_applies_response_stream_filter(self):
|
||||
graph_engine = MagicMock()
|
||||
graph_engine.run.return_value = iter([sentinel.raw_event])
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
workflow_entry.GraphEventFilterContext,
|
||||
"from_engine",
|
||||
return_value=sentinel.filter_context,
|
||||
) as from_engine,
|
||||
patch.object(
|
||||
workflow_entry,
|
||||
"ResponseStreamFilter",
|
||||
return_value=sentinel.response_stream_filter,
|
||||
) as response_stream_filter_cls,
|
||||
patch.object(
|
||||
workflow_entry,
|
||||
"filter_graph_events",
|
||||
return_value=iter([sentinel.filtered_event]),
|
||||
) as filter_graph_events,
|
||||
):
|
||||
events = list(workflow_entry.iter_dify_graph_engine_events(graph_engine))
|
||||
|
||||
assert events == [sentinel.filtered_event]
|
||||
from_engine.assert_called_once_with(graph_engine)
|
||||
response_stream_filter_cls.assert_called_once_with()
|
||||
filter_graph_events.assert_called_once_with(
|
||||
graph_engine.run.return_value,
|
||||
context=sentinel.filter_context,
|
||||
filters=[sentinel.response_stream_filter],
|
||||
)
|
||||
|
||||
def test_run_delegates_to_dify_event_iterator(self):
|
||||
entry = object.__new__(workflow_entry.WorkflowEntry)
|
||||
entry.graph_engine = sentinel.graph_engine
|
||||
|
||||
with patch.object(
|
||||
workflow_entry,
|
||||
"iter_dify_graph_engine_events",
|
||||
return_value=iter([sentinel.filtered_event]),
|
||||
) as iter_dify_graph_engine_events:
|
||||
events = list(entry.run())
|
||||
|
||||
assert events == [sentinel.filtered_event]
|
||||
iter_dify_graph_engine_events.assert_called_once_with(sentinel.graph_engine)
|
||||
|
||||
def test_run_emits_failed_event_for_unexpected_errors(self):
|
||||
entry = object.__new__(workflow_entry.WorkflowEntry)
|
||||
entry.graph_engine = MagicMock()
|
||||
|
||||
@@ -112,6 +112,29 @@ def test_enable_disable_model_load_balancing_should_call_provider_configuration_
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("method_name", "expected_provider_method"),
|
||||
[
|
||||
("enable_model_load_balancing", "enable_model_load_balancing"),
|
||||
("disable_model_load_balancing", "disable_model_load_balancing"),
|
||||
],
|
||||
)
|
||||
def test_enable_disable_model_load_balancing_uses_model_type_constructor_directly(
|
||||
method_name: str,
|
||||
expected_provider_method: str,
|
||||
service: ModelLoadBalancingService,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema())
|
||||
service.provider_manager.get_configurations.return_value = {"openai": provider_configuration}
|
||||
|
||||
getattr(service, method_name)("tenant-1", "openai", "gpt-4o-mini", "text-generation")
|
||||
|
||||
getattr(provider_configuration, expected_provider_method).assert_called_once_with(
|
||||
model="gpt-4o-mini", model_type=ModelType.LLM
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"method_name",
|
||||
["enable_model_load_balancing", "disable_model_load_balancing"],
|
||||
|
||||
@@ -368,6 +368,70 @@ class TestModelProviderServiceDelegation:
|
||||
if method_name == "get_model_credential":
|
||||
assert result == {"api_key": "x"}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("method_name", "method_kwargs", "provider_method_name", "expected_kwargs"),
|
||||
[
|
||||
(
|
||||
"get_model_credential",
|
||||
{
|
||||
"tenant_id": "tenant-1",
|
||||
"provider": "openai",
|
||||
"model_type": "text-generation",
|
||||
"model": "gpt-4o",
|
||||
"credential_id": "cred-1",
|
||||
},
|
||||
"get_custom_model_credential",
|
||||
{"model_type": ModelType.LLM, "model": "gpt-4o", "credential_id": "cred-1"},
|
||||
),
|
||||
(
|
||||
"create_model_credential",
|
||||
{
|
||||
"tenant_id": "tenant-1",
|
||||
"provider": "openai",
|
||||
"model_type": "text-generation",
|
||||
"model": "gpt-4o",
|
||||
"credentials": {"api_key": "x"},
|
||||
"credential_name": "cred-a",
|
||||
},
|
||||
"create_custom_model_credential",
|
||||
{
|
||||
"model_type": ModelType.LLM,
|
||||
"model": "gpt-4o",
|
||||
"credentials": {"api_key": "x"},
|
||||
"credential_name": "cred-a",
|
||||
},
|
||||
),
|
||||
(
|
||||
"remove_model",
|
||||
{
|
||||
"tenant_id": "tenant-1",
|
||||
"provider": "openai",
|
||||
"model_type": "text-generation",
|
||||
"model": "gpt-4o",
|
||||
},
|
||||
"delete_custom_model",
|
||||
{"model_type": ModelType.LLM, "model": "gpt-4o"},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_custom_model_methods_use_model_type_constructor_directly(
|
||||
self,
|
||||
method_name: str,
|
||||
method_kwargs: dict[str, Any],
|
||||
provider_method_name: str,
|
||||
expected_kwargs: dict[str, Any],
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
service = ModelProviderService()
|
||||
provider_configuration = MagicMock()
|
||||
get_provider_config_mock = MagicMock(return_value=provider_configuration)
|
||||
monkeypatch.setattr(service, "_get_provider_configuration", get_provider_config_mock)
|
||||
|
||||
getattr(service, method_name)(**method_kwargs)
|
||||
|
||||
get_provider_config_mock.assert_called_once_with("tenant-1", "openai")
|
||||
getattr(provider_configuration, provider_method_name).assert_called_once_with(**expected_kwargs)
|
||||
|
||||
|
||||
class TestModelProviderServiceListingsAndDefaults:
|
||||
def test_get_models_by_model_type_should_group_active_non_deprecated_models(self) -> None:
|
||||
|
||||
Reference in New Issue
Block a user