mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +08:00
refactor: replace bare dict with dict[str, Any] in model_manager and … (#35083)
This commit is contained in:
parent
3d4ddf4a6f
commit
9ed8a5ed73
@ -115,7 +115,7 @@ class ModelInstance:
|
|||||||
def invoke_llm(
|
def invoke_llm(
|
||||||
self,
|
self,
|
||||||
prompt_messages: Sequence[PromptMessage],
|
prompt_messages: Sequence[PromptMessage],
|
||||||
model_parameters: dict | None = None,
|
model_parameters: dict[str, Any] | None = None,
|
||||||
tools: Sequence[PromptMessageTool] | None = None,
|
tools: Sequence[PromptMessageTool] | None = None,
|
||||||
stop: list[str] | None = None,
|
stop: list[str] | None = None,
|
||||||
stream: Literal[True] = True,
|
stream: Literal[True] = True,
|
||||||
@ -126,7 +126,7 @@ class ModelInstance:
|
|||||||
def invoke_llm(
|
def invoke_llm(
|
||||||
self,
|
self,
|
||||||
prompt_messages: list[PromptMessage],
|
prompt_messages: list[PromptMessage],
|
||||||
model_parameters: dict | None = None,
|
model_parameters: dict[str, Any] | None = None,
|
||||||
tools: Sequence[PromptMessageTool] | None = None,
|
tools: Sequence[PromptMessageTool] | None = None,
|
||||||
stop: list[str] | None = None,
|
stop: list[str] | None = None,
|
||||||
stream: Literal[False] = False,
|
stream: Literal[False] = False,
|
||||||
@ -137,7 +137,7 @@ class ModelInstance:
|
|||||||
def invoke_llm(
|
def invoke_llm(
|
||||||
self,
|
self,
|
||||||
prompt_messages: list[PromptMessage],
|
prompt_messages: list[PromptMessage],
|
||||||
model_parameters: dict | None = None,
|
model_parameters: dict[str, Any] | None = None,
|
||||||
tools: Sequence[PromptMessageTool] | None = None,
|
tools: Sequence[PromptMessageTool] | None = None,
|
||||||
stop: list[str] | None = None,
|
stop: list[str] | None = None,
|
||||||
stream: bool = True,
|
stream: bool = True,
|
||||||
@ -147,7 +147,7 @@ class ModelInstance:
|
|||||||
def invoke_llm(
|
def invoke_llm(
|
||||||
self,
|
self,
|
||||||
prompt_messages: Sequence[PromptMessage],
|
prompt_messages: Sequence[PromptMessage],
|
||||||
model_parameters: dict | None = None,
|
model_parameters: dict[str, Any] | None = None,
|
||||||
tools: Sequence[PromptMessageTool] | None = None,
|
tools: Sequence[PromptMessageTool] | None = None,
|
||||||
stop: Sequence[str] | None = None,
|
stop: Sequence[str] | None = None,
|
||||||
stream: bool = True,
|
stream: bool = True,
|
||||||
@ -528,7 +528,7 @@ class LBModelManager:
|
|||||||
model_type: ModelType,
|
model_type: ModelType,
|
||||||
model: str,
|
model: str,
|
||||||
load_balancing_configs: list[ModelLoadBalancingConfiguration],
|
load_balancing_configs: list[ModelLoadBalancingConfiguration],
|
||||||
managed_credentials: dict | None = None,
|
managed_credentials: dict[str, Any] | None = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Load balancing model manager
|
Load balancing model manager
|
||||||
|
|||||||
@ -875,7 +875,11 @@ class DatasetRetrieval:
|
|||||||
return retrieval_resource_list
|
return retrieval_resource_list
|
||||||
|
|
||||||
def _on_retrieval_end(
|
def _on_retrieval_end(
|
||||||
self, flask_app: Flask, documents: list[Document], message_id: str | None = None, timer: dict | None = None
|
self,
|
||||||
|
flask_app: Flask,
|
||||||
|
documents: list[Document],
|
||||||
|
message_id: str | None = None,
|
||||||
|
timer: dict[str, Any] | None = None,
|
||||||
):
|
):
|
||||||
"""Handle retrieval end."""
|
"""Handle retrieval end."""
|
||||||
with flask_app.app_context():
|
with flask_app.app_context():
|
||||||
@ -980,7 +984,7 @@ class DatasetRetrieval:
|
|||||||
|
|
||||||
self._send_trace_task(message_id, documents, timer)
|
self._send_trace_task(message_id, documents, timer)
|
||||||
|
|
||||||
def _send_trace_task(self, message_id: str | None, documents: list[Document], timer: dict | None):
|
def _send_trace_task(self, message_id: str | None, documents: list[Document], timer: dict[str, Any] | None):
|
||||||
"""Send trace task if trace manager is available."""
|
"""Send trace task if trace manager is available."""
|
||||||
trace_manager: TraceQueueManager | None = (
|
trace_manager: TraceQueueManager | None = (
|
||||||
self.application_generate_entity.trace_manager if self.application_generate_entity else None
|
self.application_generate_entity.trace_manager if self.application_generate_entity else None
|
||||||
@ -1142,7 +1146,7 @@ class DatasetRetrieval:
|
|||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
hit_callback: DatasetIndexToolCallbackHandler,
|
hit_callback: DatasetIndexToolCallbackHandler,
|
||||||
user_id: str,
|
user_id: str,
|
||||||
inputs: dict,
|
inputs: dict[str, Any],
|
||||||
) -> list[DatasetRetrieverBaseTool] | None:
|
) -> list[DatasetRetrieverBaseTool] | None:
|
||||||
"""
|
"""
|
||||||
A dataset tool is a tool that can be used to retrieve information from a dataset
|
A dataset tool is a tool that can be used to retrieve information from a dataset
|
||||||
@ -1337,7 +1341,7 @@ class DatasetRetrieval:
|
|||||||
metadata_filtering_mode: str,
|
metadata_filtering_mode: str,
|
||||||
metadata_model_config: ModelConfig,
|
metadata_model_config: ModelConfig,
|
||||||
metadata_filtering_conditions: MetadataFilteringCondition | None,
|
metadata_filtering_conditions: MetadataFilteringCondition | None,
|
||||||
inputs: dict,
|
inputs: dict[str, Any],
|
||||||
) -> tuple[dict[str, list[str]] | None, MetadataFilteringCondition | None]:
|
) -> tuple[dict[str, list[str]] | None, MetadataFilteringCondition | None]:
|
||||||
document_query = select(DatasetDocument).where(
|
document_query = select(DatasetDocument).where(
|
||||||
DatasetDocument.dataset_id.in_(dataset_ids),
|
DatasetDocument.dataset_id.in_(dataset_ids),
|
||||||
@ -1417,7 +1421,7 @@ class DatasetRetrieval:
|
|||||||
metadata_filter_document_ids[document.dataset_id].append(document.id) # type: ignore
|
metadata_filter_document_ids[document.dataset_id].append(document.id) # type: ignore
|
||||||
return metadata_filter_document_ids, metadata_condition
|
return metadata_filter_document_ids, metadata_condition
|
||||||
|
|
||||||
def _replace_metadata_filter_value(self, text: str, inputs: dict) -> str:
|
def _replace_metadata_filter_value(self, text: str, inputs: dict[str, Any]) -> str:
|
||||||
if not inputs:
|
if not inputs:
|
||||||
return text
|
return text
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user