chore: add ast-grep rule to convert Optional[T] to T | None (#25560)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
-LAN- 2025-09-15 13:06:33 +08:00 • committed by GitHub
commit bab4975809
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
394 changed files with 2555 additions and 2792 deletions

View file

@ -1,5 +1,3 @@
from typing import Optional
from core.model_manager import ModelInstance, ModelManager
from core.model_runtime.entities.model_entities import ModelType
from core.model_runtime.errors.invoke import InvokeAuthorizationError
@ -18,8 +16,8 @@ class DataPostProcessor:
self,
tenant_id: str,
reranking_mode: str,
reranking_model: Optional[dict] = None,
weights: Optional[dict] = None,
reranking_model: dict | None = None,
weights: dict | None = None,
reorder_enabled: bool = False,
):
self.rerank_runner = self._get_rerank_runner(reranking_mode, tenant_id, reranking_model, weights)
@ -29,9 +27,9 @@ class DataPostProcessor:
self,
query: str,
documents: list[Document],
score_threshold: Optional[float] = None,
top_n: Optional[int] = None,
user: Optional[str] = None,
score_threshold: float | None = None,
top_n: int | None = None,
user: str | None = None,
) -> list[Document]:
if self.rerank_runner:
documents = self.rerank_runner.run(query, documents, score_threshold, top_n, user)
@ -45,9 +43,9 @@ class DataPostProcessor:
self,
reranking_mode: str,
tenant_id: str,
reranking_model: Optional[dict] = None,
weights: Optional[dict] = None,
) -> Optional[BaseRerankRunner]:
reranking_model: dict | None = None,
weights: dict | None = None,
) -> BaseRerankRunner | None:
if reranking_mode == RerankMode.WEIGHTED_SCORE.value and weights:
runner = RerankRunnerFactory.create_rerank_runner(
runner_type=reranking_mode,
@ -74,12 +72,12 @@ class DataPostProcessor:
return runner
return None
def _get_reorder_runner(self, reorder_enabled) -> Optional[ReorderRunner]:
def _get_reorder_runner(self, reorder_enabled) -> ReorderRunner | None:
if reorder_enabled:
return ReorderRunner()
return None
def _get_rerank_model_instance(self, tenant_id: str, reranking_model: Optional[dict]) -> ModelInstance | None:
def _get_rerank_model_instance(self, tenant_id: str, reranking_model: dict | None) -> ModelInstance | None:
if reranking_model:
try:
model_manager = ModelManager()

View file

@ -1,5 +1,5 @@
from collections import defaultdict
from typing import Any, Optional
from typing import Any
import orjson
from pydantic import BaseModel
@ -143,7 +143,7 @@ class Jieba(BaseKeyword):
storage.delete(file_key)
storage.save(file_key, dumps_with_sets(keyword_table_dict).encode("utf-8"))
def _get_dataset_keyword_table(self) -> Optional[dict]:
def _get_dataset_keyword_table(self) -> dict | None:
dataset_keyword_table = self.dataset.dataset_keyword_table
if dataset_keyword_table:
keyword_table_dict = dataset_keyword_table.keyword_table_dict

View file

@ -1,5 +1,5 @@
import re
from typing import Optional, cast
from typing import cast
class JiebaKeywordTableHandler:
@ -10,7 +10,7 @@ class JiebaKeywordTableHandler:
jieba.analyse.default_tfidf.stop_words = STOPWORDS # type: ignore
def extract_keywords(self, text: str, max_keywords_per_chunk: Optional[int] = 10) -> set[str]:
def extract_keywords(self, text: str, max_keywords_per_chunk: int | None = 10) -> set[str]:
"""Extract keywords with JIEBA tfidf."""
import jieba.analyse # type: ignore

View file

@ -1,6 +1,5 @@
import concurrent.futures
from concurrent.futures import ThreadPoolExecutor
from typing import Optional
from flask import Flask, current_app
from sqlalchemy import select
@ -39,11 +38,11 @@ class RetrievalService:
dataset_id: str,
query: str,
top_k: int,
score_threshold: Optional[float] = 0.0,
reranking_model: Optional[dict] = None,
score_threshold: float | None = 0.0,
reranking_model: dict | None = None,
reranking_mode: str = "reranking_model",
weights: Optional[dict] = None,
document_ids_filter: Optional[list[str]] = None,
weights: dict | None = None,
document_ids_filter: list[str] | None = None,
):
if not query:
return []
@ -125,8 +124,8 @@ class RetrievalService:
cls,
dataset_id: str,
query: str,
external_retrieval_model: Optional[dict] = None,
metadata_filtering_conditions: Optional[dict] = None,
external_retrieval_model: dict | None = None,
metadata_filtering_conditions: dict | None = None,
):
stmt = select(Dataset).where(Dataset.id == dataset_id)
dataset = db.session.scalar(stmt)
@ -145,7 +144,7 @@ class RetrievalService:
return all_documents
@classmethod
def _get_dataset(cls, dataset_id: str) -> Optional[Dataset]:
def _get_dataset(cls, dataset_id: str) -> Dataset | None:
with Session(db.engine) as session:
return session.query(Dataset).where(Dataset.id == dataset_id).first()
@ -158,7 +157,7 @@ class RetrievalService:
top_k: int,
all_documents: list,
exceptions: list,
document_ids_filter: Optional[list[str]] = None,
document_ids_filter: list[str] | None = None,
):
with flask_app.app_context():
try:
@ -182,12 +181,12 @@ class RetrievalService:
dataset_id: str,
query: str,
top_k: int,
score_threshold: Optional[float],
reranking_model: Optional[dict],
score_threshold: float | None,
reranking_model: dict | None,
all_documents: list,
retrieval_method: str,
exceptions: list,
document_ids_filter: Optional[list[str]] = None,
document_ids_filter: list[str] | None = None,
):
with flask_app.app_context():
try:
@ -235,12 +234,12 @@ class RetrievalService:
dataset_id: str,
query: str,
top_k: int,
score_threshold: Optional[float],
reranking_model: Optional[dict],
score_threshold: float | None,
reranking_model: dict | None,
all_documents: list,
retrieval_method: str,
exceptions: list,
document_ids_filter: Optional[list[str]] = None,
document_ids_filter: list[str] | None = None,
):
with flask_app.app_context():
try:

View file

@ -1,5 +1,5 @@
import json
from typing import Any, Optional
from typing import Any
from pydantic import BaseModel, model_validator
@ -20,7 +20,7 @@ class AnalyticdbVectorOpenAPIConfig(BaseModel):
account: str
account_password: str
namespace: str = "dify"
namespace_password: Optional[str] = None
namespace_password: str | None = None
metrics: str = "cosine"
read_timeout: int = 60000

View file

@ -1,5 +1,5 @@
import json
from typing import Any, Optional
from typing import Any
import chromadb
from chromadb import QueryResult, Settings
@ -20,8 +20,8 @@ class ChromaConfig(BaseModel):
port: int
tenant: str
database: str
auth_provider: Optional[str] = None
auth_credentials: Optional[str] = None
auth_provider: str | None = None
auth_credentials: str | None = None
def to_chroma_params(self):
settings = Settings(

View file

@ -84,7 +84,7 @@ class ClickzettaConnectionPool:
self._pool_locks: dict[str, threading.Lock] = {}
self._max_pool_size = 5 # Maximum connections per configuration
self._connection_timeout = 300 # 5 minutes timeout
self._cleanup_thread: Optional[threading.Thread] = None
self._cleanup_thread: threading.Thread | None = None
self._shutdown = False
self._start_cleanup_thread()
@ -303,8 +303,8 @@ class ClickzettaVector(BaseVector):
"""
# Class-level write queue and lock for serializing writes
_write_queue: Optional[queue.Queue] = None
_write_thread: Optional[threading.Thread] = None
_write_queue: queue.Queue | None = None
_write_thread: threading.Thread | None = None
_write_lock = threading.Lock()
_shutdown = False
@ -328,7 +328,7 @@ class ClickzettaVector(BaseVector):
def __init__(self, vector_instance: "ClickzettaVector"):
self.vector = vector_instance
self.connection: Optional[Connection] = None
self.connection: Connection | None = None
def __enter__(self) -> "Connection":
self.connection = self.vector._get_connection()

View file

@ -1,6 +1,6 @@
import json
import logging
from typing import Any, Optional
from typing import Any
from flask import current_app
@ -22,8 +22,8 @@ class ElasticSearchJaVector(ElasticSearchVector):
def create_collection(
self,
embeddings: list[list[float]],
metadatas: Optional[list[dict[Any, Any]]] = None,
index_params: Optional[dict] = None,
metadatas: list[dict[Any, Any]] | None = None,
index_params: dict | None = None,
):
lock_name = f"vector_indexing_lock_{self._collection_name}"
with redis_client.lock(lock_name, timeout=20):

View file

@ -1,7 +1,7 @@
import json
import logging
import math
from typing import Any, Optional, cast
from typing import Any, cast
from urllib.parse import urlparse
import requests
@ -24,18 +24,18 @@ logger = logging.getLogger(__name__)
class ElasticSearchConfig(BaseModel):
# Regular Elasticsearch config
host: Optional[str] = None
port: Optional[int] = None
username: Optional[str] = None
password: Optional[str] = None
host: str | None = None
port: int | None = None
username: str | None = None
password: str | None = None
# Elastic Cloud specific config
cloud_url: Optional[str] = None # Cloud URL for Elasticsearch Cloud
api_key: Optional[str] = None
cloud_url: str | None = None # Cloud URL for Elasticsearch Cloud
api_key: str | None = None
# Common config
use_cloud: bool = False
ca_certs: Optional[str] = None
ca_certs: str | None = None
verify_certs: bool = False
request_timeout: int = 100000
retry_on_timeout: bool = True
@ -256,8 +256,8 @@ class ElasticSearchVector(BaseVector):
def create_collection(
self,
embeddings: list[list[float]],
metadatas: Optional[list[dict[Any, Any]]] = None,
index_params: Optional[dict] = None,
metadatas: list[dict[Any, Any]] | None = None,
index_params: dict | None = None,
):
lock_name = f"vector_indexing_lock_{self._collection_name}"
with redis_client.lock(lock_name, timeout=20):

View file

@ -1,7 +1,7 @@
import json
import logging
import ssl
from typing import Any, Optional
from typing import Any
from elasticsearch import Elasticsearch
from pydantic import BaseModel, model_validator
@ -157,8 +157,8 @@ class HuaweiCloudVector(BaseVector):
def create_collection(
self,
embeddings: list[list[float]],
metadatas: Optional[list[dict[Any, Any]]] = None,
index_params: Optional[dict] = None,
metadatas: list[dict[Any, Any]] | None = None,
index_params: dict | None = None,
):
lock_name = f"vector_indexing_lock_{self._collection_name}"
with redis_client.lock(lock_name, timeout=20):

View file

@ -2,7 +2,7 @@ import copy
import json
import logging
import time
from typing import Any, Optional
from typing import Any
from opensearchpy import OpenSearch, helpers
from opensearchpy.helpers import BulkIndexError
@ -29,10 +29,10 @@ UGC_INDEX_PREFIX = "ugc_index"
class LindormVectorStoreConfig(BaseModel):
hosts: str
username: Optional[str] = None
password: Optional[str] = None
using_ugc: Optional[bool] = False
request_timeout: Optional[float] = 1.0 # timeout units: s
username: str | None = None
password: str | None = None
using_ugc: bool | None = False
request_timeout: float | None = 1.0 # timeout units: s
@model_validator(mode="before")
@classmethod
@ -448,13 +448,13 @@ def default_text_search_query(
query_text: str,
k: int = 4,
text_field: str = Field.CONTENT_KEY.value,
must: Optional[list[dict]] = None,
must_not: Optional[list[dict]] = None,
should: Optional[list[dict]] = None,
must: list[dict] | None = None,
must_not: list[dict] | None = None,
should: list[dict] | None = None,
minimum_should_match: int = 0,
filters: Optional[list[dict]] = None,
routing: Optional[str] = None,
routing_field: Optional[str] = None,
filters: list[dict] | None = None,
routing: str | None = None,
routing_field: str | None = None,
**kwargs,
):
query_clause: dict[str, Any] = {}
@ -505,13 +505,13 @@ def default_vector_search_query(
query_vector: list[float],
k: int = 4,
min_score: str = "0.0",
ef_search: Optional[str] = None, # only for hnsw
nprobe: Optional[str] = None, # "2000"
reorder_factor: Optional[str] = None, # "20"
client_refactor: Optional[str] = None, # "true"
ef_search: str | None = None, # only for hnsw
nprobe: str | None = None, # "2000"
reorder_factor: str | None = None, # "20"
client_refactor: str | None = None, # "true"
vector_field: str = Field.VECTOR.value,
filters: Optional[list[dict]] = None,
filter_type: Optional[str] = None,
filters: list[dict] | None = None,
filter_type: str | None = None,
**kwargs,
):
if filters is not None:

View file

@ -3,7 +3,7 @@ import logging
import uuid
from collections.abc import Callable
from functools import wraps
from typing import Any, Concatenate, Optional, ParamSpec, TypeVar
from typing import Any, Concatenate, ParamSpec, TypeVar
from mo_vector.client import MoVectorClient # type: ignore
from pydantic import BaseModel, model_validator
@ -74,7 +74,7 @@ class MatrixoneVector(BaseVector):
self.client = self._get_client(len(embeddings[0]), True)
return self.add_texts(texts, embeddings)
def _get_client(self, dimension: Optional[int] = None, create_table: bool = False) -> MoVectorClient:
def _get_client(self, dimension: int | None = None, create_table: bool = False) -> MoVectorClient:
"""
Create a new client for the collection.

View file

@ -1,6 +1,6 @@
import json
import logging
from typing import Any, Optional
from typing import Any
from packaging import version
from pydantic import BaseModel, model_validator
@ -26,13 +26,13 @@ class MilvusConfig(BaseModel):
"""
uri: str # Milvus server URI
token: Optional[str] = None # Optional token for authentication
user: Optional[str] = None # Username for authentication
password: Optional[str] = None # Password for authentication
token: str | None = None # Optional token for authentication
user: str | None = None # Username for authentication
password: str | None = None # Password for authentication
batch_size: int = 100 # Batch size for operations
database: str = "default" # Database name
enable_hybrid_search: bool = False # Flag to enable hybrid search
analyzer_params: Optional[str] = None # Analyzer params
analyzer_params: str | None = None # Analyzer params
@model_validator(mode="before")
@classmethod
@ -79,7 +79,7 @@ class MilvusVector(BaseVector):
self._load_collection_fields()
self._hybrid_search_enabled = self._check_hybrid_search_support() # Check if hybrid search is supported
def _load_collection_fields(self, fields: Optional[list[str]] = None):
def _load_collection_fields(self, fields: list[str] | None = None):
if fields is None:
# Load collection fields from remote server
collection_info = self._client.describe_collection(self._collection_name)
@ -292,7 +292,7 @@ class MilvusVector(BaseVector):
)
def create_collection(
self, embeddings: list, metadatas: Optional[list[dict]] = None, index_params: Optional[dict] = None
self, embeddings: list, metadatas: list[dict] | None = None, index_params: dict | None = None
):
"""
Create a new collection in Milvus with the specified schema and index parameters.

View file

@ -1,6 +1,6 @@
import json
import logging
from typing import Any, Literal, Optional
from typing import Any, Literal
from uuid import uuid4
from opensearchpy import OpenSearch, Urllib3AWSV4SignerAuth, Urllib3HttpConnection, helpers
@ -26,10 +26,10 @@ class OpenSearchConfig(BaseModel):
secure: bool = False # use_ssl
verify_certs: bool = True
auth_method: Literal["basic", "aws_managed_iam"] = "basic"
user: Optional[str] = None
password: Optional[str] = None
aws_region: Optional[str] = None
aws_service: Optional[str] = None
user: str | None = None
password: str | None = None
aws_region: str | None = None
aws_service: str | None = None
@model_validator(mode="before")
@classmethod
@ -236,7 +236,7 @@ class OpenSearchVector(BaseVector):
return docs
def create_collection(
self, embeddings: list, metadatas: Optional[list[dict]] = None, index_params: Optional[dict] = None
self, embeddings: list, metadatas: list[dict] | None = None, index_params: dict | None = None
):
lock_name = f"vector_indexing_lock_{self._collection_name.lower()}"
with redis_client.lock(lock_name, timeout=20):

View file

@ -3,7 +3,7 @@ import os
import uuid
from collections.abc import Generator, Iterable, Sequence
from itertools import islice
from typing import TYPE_CHECKING, Any, Optional, Union
from typing import TYPE_CHECKING, Any, Union
import qdrant_client
from flask import current_app
@ -46,7 +46,7 @@ class PathQdrantParams(BaseModel):
class UrlQdrantParams(BaseModel):
url: str
api_key: Optional[str]
api_key: str | None
timeout: float
verify: bool
grpc_port: int
@ -55,9 +55,9 @@ class UrlQdrantParams(BaseModel):
class QdrantConfig(BaseModel):
endpoint: str
api_key: Optional[str] = None
api_key: str | None = None
timeout: float = 20
root_path: Optional[str] = None
root_path: str | None = None
grpc_port: int = 6334
prefer_grpc: bool = False
replication_factor: int = 1
@ -189,10 +189,10 @@ class QdrantVector(BaseVector):
self,
texts: Iterable[str],
embeddings: list[list[float]],
metadatas: Optional[list[dict]] = None,
ids: Optional[Sequence[str]] = None,
metadatas: list[dict] | None = None,
ids: Sequence[str] | None = None,
batch_size: int = 64,
group_id: Optional[str] = None,
group_id: str | None = None,
) -> Generator[tuple[list[str], list[rest.PointStruct]], None, None]:
from qdrant_client.http import models as rest
@ -234,7 +234,7 @@ class QdrantVector(BaseVector):
def _build_payloads(
cls,
texts: Iterable[str],
metadatas: Optional[list[dict]],
metadatas: list[dict] | None,
content_payload_key: str,
metadata_payload_key: str,
group_id: str,

View file

@ -1,6 +1,6 @@
import json
import uuid
from typing import Any, Optional
from typing import Any
from pydantic import BaseModel, model_validator
from sqlalchemy import Column, String, Table, create_engine, insert
@ -160,7 +160,7 @@ class RelytVector(BaseVector):
else:
return None
def delete_by_uuids(self, ids: Optional[list[str]] = None):
def delete_by_uuids(self, ids: list[str] | None = None):
"""Delete by vector IDs.
Args:
@ -241,7 +241,7 @@ class RelytVector(BaseVector):
self,
embedding: list[float],
k: int = 4,
filter: Optional[dict] = None,
filter: dict | None = None,
) -> list[tuple[Document, float]]:
# Add the filter if provided

View file

@ -2,7 +2,7 @@ import json
import logging
import math
from collections.abc import Iterable
from typing import Any, Optional
from typing import Any
import tablestore # type: ignore
from pydantic import BaseModel, model_validator
@ -22,11 +22,11 @@ logger = logging.getLogger(__name__)
class TableStoreConfig(BaseModel):
access_key_id: Optional[str] = None
access_key_secret: Optional[str] = None
instance_name: Optional[str] = None
endpoint: Optional[str] = None
normalize_full_text_bm25_score: Optional[bool] = False
access_key_id: str | None = None
access_key_secret: str | None = None
instance_name: str | None = None
endpoint: str | None = None
normalize_full_text_bm25_score: bool | None = False
@model_validator(mode="before")
@classmethod

View file

@ -1,7 +1,7 @@
import json
import logging
import math
from typing import Any, Optional
from typing import Any
from pydantic import BaseModel
from tcvdb_text.encoder import BM25Encoder # type: ignore
@ -24,10 +24,10 @@ logger = logging.getLogger(__name__)
class TencentConfig(BaseModel):
url: str
api_key: Optional[str] = None
api_key: str | None = None
timeout: float = 30
username: Optional[str] = None
database: Optional[str] = None
username: str | None = None
database: str | None = None
index_type: str = "HNSW"
metric_type: str = "IP"
shard: int = 1

View file

@ -3,7 +3,7 @@ import os
import uuid
from collections.abc import Generator, Iterable, Sequence
from itertools import islice
from typing import TYPE_CHECKING, Any, Optional, Union
from typing import TYPE_CHECKING, Any, Union
import qdrant_client
import requests
@ -45,9 +45,9 @@ if TYPE_CHECKING:
class TidbOnQdrantConfig(BaseModel):
endpoint: str
api_key: Optional[str] = None
api_key: str | None = None
timeout: float = 20
root_path: Optional[str] = None
root_path: str | None = None
grpc_port: int = 6334
prefer_grpc: bool = False
replication_factor: int = 1
@ -180,10 +180,10 @@ class TidbOnQdrantVector(BaseVector):
self,
texts: Iterable[str],
embeddings: list[list[float]],
metadatas: Optional[list[dict]] = None,
ids: Optional[Sequence[str]] = None,
metadatas: list[dict] | None = None,
ids: Sequence[str] | None = None,
batch_size: int = 64,
group_id: Optional[str] = None,
group_id: str | None = None,
) -> Generator[tuple[list[str], list[rest.PointStruct]], None, None]:
from qdrant_client.http import models as rest
@ -225,7 +225,7 @@ class TidbOnQdrantVector(BaseVector):
def _build_payloads(
cls,
texts: Iterable[str],
metadatas: Optional[list[dict]],
metadatas: list[dict] | None,
content_payload_key: str,
metadata_payload_key: str,
group_id: str,

View file

@ -1,7 +1,7 @@
import logging
import time
from abc import ABC, abstractmethod
from typing import Any, Optional
from typing import Any
from sqlalchemy import select
@ -32,7 +32,7 @@ class AbstractVectorFactory(ABC):
class Vector:
def __init__(self, dataset: Dataset, attributes: Optional[list] = None):
def __init__(self, dataset: Dataset, attributes: list | None = None):
if attributes is None:
attributes = ["doc_id", "dataset_id", "document_id", "doc_hash"]
self._dataset = dataset
@ -180,7 +180,7 @@ class Vector:
case _:
raise ValueError(f"Vector store {vector_type} is not supported.")
def create(self, texts: Optional[list] = None, **kwargs):
def create(self, texts: list | None = None, **kwargs):
if texts:
start = time.time()
logger.info("start embedding %s texts %s", len(texts), start)

View file

@ -1,6 +1,6 @@
import datetime
import json
from typing import Any, Optional
from typing import Any
import requests
import weaviate # type: ignore
@ -19,7 +19,7 @@ from models.dataset import Dataset
class WeaviateConfig(BaseModel):
endpoint: str
api_key: Optional[str] = None
api_key: str | None = None
batch_size: int = 100
@model_validator(mode="before")

View file

@ -1,5 +1,5 @@
from collections.abc import Sequence
from typing import Any, Optional
from typing import Any
from sqlalchemy import func, select
@ -15,7 +15,7 @@ class DatasetDocumentStore:
self,
dataset: Dataset,
user_id: str,
document_id: Optional[str] = None,
document_id: str | None = None,
):
self._dataset = dataset
self._user_id = user_id
@ -176,7 +176,7 @@ class DatasetDocumentStore:
result = self.get_document_segment(doc_id)
return result is not None
def get_document(self, doc_id: str, raise_error: bool = True) -> Optional[Document]:
def get_document(self, doc_id: str, raise_error: bool = True) -> Document | None:
document_segment = self.get_document_segment(doc_id)
if document_segment is None:
@ -217,16 +217,16 @@ class DatasetDocumentStore:
document_segment.index_node_hash = doc_hash
db.session.commit()
def get_document_hash(self, doc_id: str) -> Optional[str]:
def get_document_hash(self, doc_id: str) -> str | None:
"""Get the stored hash for a document, if it exists."""
document_segment = self.get_document_segment(doc_id)
if document_segment is None:
return None
data: Optional[str] = document_segment.index_node_hash
data: str | None = document_segment.index_node_hash
return data
def get_document_segment(self, doc_id: str) -> Optional[DocumentSegment]:
def get_document_segment(self, doc_id: str) -> DocumentSegment | None:
stmt = select(DocumentSegment).where(
DocumentSegment.dataset_id == self._dataset.id, DocumentSegment.index_node_id == doc_id
)

View file

@ -1,6 +1,6 @@
import base64
import logging
from typing import Any, Optional, cast
from typing import Any, cast
import numpy as np
from sqlalchemy.exc import IntegrityError
@ -20,7 +20,7 @@ logger = logging.getLogger(__name__)
class CacheEmbedding(Embeddings):
def __init__(self, model_instance: ModelInstance, user: Optional[str] = None):
def __init__(self, model_instance: ModelInstance, user: str | None = None):
self._model_instance = model_instance
self._user = user

View file

@ -1,5 +1,3 @@
from typing import Optional
from pydantic import BaseModel
from models.dataset import DocumentSegment
@ -19,5 +17,5 @@ class RetrievalSegments(BaseModel):
model_config = {"arbitrary_types_allowed": True}
segment: DocumentSegment
child_chunks: Optional[list[RetrievalChildChunk]] = None
score: Optional[float] = None
child_chunks: list[RetrievalChildChunk] | None = None
score: float | None = None

View file

@ -1,23 +1,23 @@
from typing import Any, Optional
from typing import Any
from pydantic import BaseModel
class RetrievalSourceMetadata(BaseModel):
position: Optional[int] = None
dataset_id: Optional[str] = None
dataset_name: Optional[str] = None
document_id: Optional[str] = None
document_name: Optional[str] = None
data_source_type: Optional[str] = None
segment_id: Optional[str] = None
retriever_from: Optional[str] = None
score: Optional[float] = None
hit_count: Optional[int] = None
word_count: Optional[int] = None
segment_position: Optional[int] = None
index_node_hash: Optional[str] = None
content: Optional[str] = None
page: Optional[int] = None
doc_metadata: Optional[dict[str, Any]] = None
title: Optional[str] = None
position: int | None = None
dataset_id: str | None = None
dataset_name: str | None = None
document_id: str | None = None
document_name: str | None = None
data_source_type: str | None = None
segment_id: str | None = None
retriever_from: str | None = None
score: float | None = None
hit_count: int | None = None
word_count: int | None = None
segment_position: int | None = None
index_node_hash: str | None = None
content: str | None = None
page: int | None = None
doc_metadata: dict[str, Any] | None = None
title: str | None = None

View file

@ -1,5 +1,3 @@
from typing import Optional
from pydantic import BaseModel
@ -9,4 +7,4 @@ class DocumentContext(BaseModel):
"""
content: str
score: Optional[float] = None
score: float | None = None

View file

@ -1,5 +1,5 @@
from collections.abc import Sequence
from typing import Literal, Optional
from typing import Literal
from pydantic import BaseModel, Field
@ -43,5 +43,5 @@ class MetadataCondition(BaseModel):
Metadata Condition.
"""
logical_operator: Optional[Literal["and", "or"]] = "and"
conditions: Optional[list[Condition]] = Field(default=None, deprecated=True)
logical_operator: Literal["and", "or"] | None = "and"
conditions: list[Condition] | None = Field(default=None, deprecated=True)

View file

@ -12,7 +12,7 @@ import mimetypes
from collections.abc import Generator, Mapping
from io import BufferedReader, BytesIO
from pathlib import Path, PurePath
from typing import Any, Optional, Union
from typing import Any, Union
from pydantic import BaseModel, ConfigDict, model_validator
@ -30,17 +30,17 @@ class Blob(BaseModel):
"""
data: Union[bytes, str, None] = None # Raw data
mimetype: Optional[str] = None # Not to be confused with a file extension
mimetype: str | None = None # Not to be confused with a file extension
encoding: str = "utf-8" # Use utf-8 as default encoding, if decoding to string
# Location where the original content was found
# Represent location on the local file system
# Useful for situations where downstream code assumes it must work with file paths
# rather than in-memory content.
path: Optional[PathLike] = None
path: PathLike | None = None
model_config = ConfigDict(arbitrary_types_allowed=True, frozen=True)
@property
def source(self) -> Optional[str]:
def source(self) -> str | None:
"""The source location of the blob as string if known otherwise none."""
return str(self.path) if self.path else None
@ -91,7 +91,7 @@ class Blob(BaseModel):
path: PathLike,
*,
encoding: str = "utf-8",
mime_type: Optional[str] = None,
mime_type: str | None = None,
guess_type: bool = True,
) -> Blob:
"""Load the blob from a path like object.
@ -120,8 +120,8 @@ class Blob(BaseModel):
data: Union[str, bytes],
*,
encoding: str = "utf-8",
mime_type: Optional[str] = None,
path: Optional[str] = None,
mime_type: str | None = None,
path: str | None = None,
) -> Blob:
"""Initialize the blob from in-memory data.

View file

@ -1,7 +1,6 @@
"""Abstract interface for document loader implementations."""
import csv
from typing import Optional
import pandas as pd
@ -21,10 +20,10 @@ class CSVExtractor(BaseExtractor):
def __init__(
self,
file_path: str,
encoding: Optional[str] = None,
encoding: str | None = None,
autodetect_encoding: bool = False,
source_column: Optional[str] = None,
csv_args: Optional[dict] = None,
source_column: str | None = None,
csv_args: dict | None = None,
):
"""Initialize with file path."""
self._file_path = file_path

View file

@ -1,5 +1,3 @@
from typing import Optional
from pydantic import BaseModel, ConfigDict
from models.dataset import Document
@ -14,7 +12,7 @@ class NotionInfo(BaseModel):
notion_workspace_id: str
notion_obj_id: str
notion_page_type: str
document: Optional[Document] = None
document: Document | None = None
tenant_id: str
model_config = ConfigDict(arbitrary_types_allowed=True)
@ -43,10 +41,10 @@ class ExtractSetting(BaseModel):
"""
datasource_type: str
upload_file: Optional[UploadFile] = None
notion_info: Optional[NotionInfo] = None
website_info: Optional[WebsiteInfo] = None
document_model: Optional[str] = None
upload_file: UploadFile | None = None
notion_info: NotionInfo | None = None
website_info: WebsiteInfo | None = None
document_model: str | None = None
model_config = ConfigDict(arbitrary_types_allowed=True)
def __init__(self, **data):

View file

@ -1,7 +1,7 @@
"""Abstract interface for document loader implementations."""
import os
from typing import Optional, cast
from typing import cast
import pandas as pd
from openpyxl import load_workbook
@ -18,7 +18,7 @@ class ExcelExtractor(BaseExtractor):
file_path: Path to the file to load.
"""
def __init__(self, file_path: str, encoding: Optional[str] = None, autodetect_encoding: bool = False):
def __init__(self, file_path: str, encoding: str | None = None, autodetect_encoding: bool = False):
"""Initialize with file path."""
self._file_path = file_path
self._encoding = encoding

View file

@ -1,7 +1,7 @@
import re
import tempfile
from pathlib import Path
from typing import Optional, Union
from typing import Union
from urllib.parse import unquote
from configs import dify_config
@ -90,7 +90,7 @@ class ExtractProcessor:
@classmethod
def extract(
cls, extract_setting: ExtractSetting, is_automatic: bool = False, file_path: Optional[str] = None
cls, extract_setting: ExtractSetting, is_automatic: bool = False, file_path: str | None = None
) -> list[Document]:
if extract_setting.datasource_type == DatasourceType.FILE.value:
with tempfile.TemporaryDirectory() as temp_dir:
@ -104,7 +104,7 @@ class ExtractProcessor:
input_file = Path(file_path)
file_extension = input_file.suffix.lower()
etl_type = dify_config.ETL_TYPE
extractor: Optional[BaseExtractor] = None
extractor: BaseExtractor | None = None
if etl_type == "Unstructured":
unstructured_api_url = dify_config.UNSTRUCTURED_API_URL or ""
unstructured_api_key = dify_config.UNSTRUCTURED_API_KEY or ""

View file

@ -1,17 +1,17 @@
"""Document loader helpers."""
import concurrent.futures
from typing import NamedTuple, Optional, cast
from typing import NamedTuple, cast
class FileEncoding(NamedTuple):
"""A file encoding as the NamedTuple."""
encoding: Optional[str]
encoding: str | None
"""The encoding of the file."""
confidence: float
"""The confidence of the encoding."""
language: Optional[str]
language: str | None
"""The language of the file."""

View file

@ -2,7 +2,6 @@
import re
from pathlib import Path
from typing import Optional
from core.rag.extractor.extractor_base import BaseExtractor
from core.rag.extractor.helpers import detect_file_encodings
@ -22,7 +21,7 @@ class MarkdownExtractor(BaseExtractor):
file_path: str,
remove_hyperlinks: bool = False,
remove_images: bool = False,
encoding: Optional[str] = None,
encoding: str | None = None,
autodetect_encoding: bool = True,
):
"""Initialize with file path."""
@ -45,13 +44,13 @@ class MarkdownExtractor(BaseExtractor):
return documents
def markdown_to_tups(self, markdown_text: str) -> list[tuple[Optional[str], str]]:
def markdown_to_tups(self, markdown_text: str) -> list[tuple[str | None, str]]:
"""Convert a markdown file to a dictionary.
The keys are the headers and the values are the text under each header.
"""
markdown_tups: list[tuple[Optional[str], str]] = []
markdown_tups: list[tuple[str | None, str]] = []
lines = markdown_text.split("\n")
current_header = None
@ -94,7 +93,7 @@ class MarkdownExtractor(BaseExtractor):
content = re.sub(pattern, r"\1", content)
return content
def parse_tups(self, filepath: str) -> list[tuple[Optional[str], str]]:
def parse_tups(self, filepath: str) -> list[tuple[str | None, str]]:
"""Parse file into tuples."""
content = ""
try:

View file

@ -1,7 +1,7 @@
import json
import logging
import operator
from typing import Any, Optional, cast
from typing import Any, cast
import requests
from sqlalchemy import select
@ -36,8 +36,8 @@ class NotionExtractor(BaseExtractor):
notion_obj_id: str,
notion_page_type: str,
tenant_id: str,
document_model: Optional[DocumentModel] = None,
notion_access_token: Optional[str] = None,
document_model: DocumentModel | None = None,
notion_access_token: str | None = None,
):
self._notion_access_token = None
self._document_model = document_model
@ -328,7 +328,7 @@ class NotionExtractor(BaseExtractor):
result_lines = "\n".join(result_lines_arr)
return result_lines
def update_last_edited_time(self, document_model: Optional[DocumentModel]):
def update_last_edited_time(self, document_model: DocumentModel | None):
if not document_model:
return

View file

@ -2,7 +2,6 @@
import contextlib
from collections.abc import Iterator
from typing import Optional
from core.rag.extractor.blob.blob import Blob
from core.rag.extractor.extractor_base import BaseExtractor
@ -18,7 +17,7 @@ class PdfExtractor(BaseExtractor):
file_path: Path to the file to load.
"""
def __init__(self, file_path: str, file_cache_key: Optional[str] = None):
def __init__(self, file_path: str, file_cache_key: str | None = None):
"""Initialize with file path."""
self._file_path = file_path
self._file_cache_key = file_cache_key

View file

@ -1,7 +1,6 @@
"""Abstract interface for document loader implementations."""
from pathlib import Path
from typing import Optional
from core.rag.extractor.extractor_base import BaseExtractor
from core.rag.extractor.helpers import detect_file_encodings
@ -16,7 +15,7 @@ class TextExtractor(BaseExtractor):
file_path: Path to the file to load.
"""
def __init__(self, file_path: str, encoding: Optional[str] = None, autodetect_encoding: bool = False):
def __init__(self, file_path: str, encoding: str | None = None, autodetect_encoding: bool = False):
"""Initialize with file path."""
self._file_path = file_path
self._encoding = encoding

View file

@ -1,7 +1,6 @@
import base64
import contextlib
import logging
from typing import Optional
from bs4 import BeautifulSoup
@ -17,7 +16,7 @@ class UnstructuredEmailExtractor(BaseExtractor):
file_path: Path to the file to load.
"""
def __init__(self, file_path: str, api_url: Optional[str] = None, api_key: str = ""):
def __init__(self, file_path: str, api_url: str | None = None, api_key: str = ""):
"""Initialize with file path."""
self._file_path = file_path
self._api_url = api_url

View file

@ -1,5 +1,4 @@
import logging
from typing import Optional
import pypandoc # type: ignore
@ -20,7 +19,7 @@ class UnstructuredEpubExtractor(BaseExtractor):
def __init__(
self,
file_path: str,
api_url: Optional[str] = None,
api_url: str | None = None,
api_key: str = "",
):
"""Initialize with file path."""

View file

@ -1,5 +1,4 @@
import logging
from typing import Optional
from core.rag.extractor.extractor_base import BaseExtractor
from core.rag.models.document import Document
@ -16,7 +15,7 @@ class UnstructuredMarkdownExtractor(BaseExtractor):
"""
def __init__(self, file_path: str, api_url: Optional[str] = None, api_key: str = ""):
def __init__(self, file_path: str, api_url: str | None = None, api_key: str = ""):
"""Initialize with file path."""
self._file_path = file_path
self._api_url = api_url

View file

@ -1,5 +1,4 @@
import logging
from typing import Optional
from core.rag.extractor.extractor_base import BaseExtractor
from core.rag.models.document import Document
@ -15,7 +14,7 @@ class UnstructuredMsgExtractor(BaseExtractor):
file_path: Path to the file to load.
"""
def __init__(self, file_path: str, api_url: Optional[str] = None, api_key: str = ""):
def __init__(self, file_path: str, api_url: str | None = None, api_key: str = ""):
"""Initialize with file path."""
self._file_path = file_path
self._api_url = api_url

View file

@ -1,5 +1,4 @@
import logging
from typing import Optional
from core.rag.extractor.extractor_base import BaseExtractor
from core.rag.models.document import Document
@ -15,7 +14,7 @@ class UnstructuredPPTExtractor(BaseExtractor):
file_path: Path to the file to load.
"""
def __init__(self, file_path: str, api_url: Optional[str] = None, api_key: str = ""):
def __init__(self, file_path: str, api_url: str | None = None, api_key: str = ""):
"""Initialize with file path."""
self._file_path = file_path
self._api_url = api_url

View file

@ -1,5 +1,4 @@
import logging
from typing import Optional
from core.rag.extractor.extractor_base import BaseExtractor
from core.rag.models.document import Document
@ -15,7 +14,7 @@ class UnstructuredPPTXExtractor(BaseExtractor):
file_path: Path to the file to load.
"""
def __init__(self, file_path: str, api_url: Optional[str] = None, api_key: str = ""):
def __init__(self, file_path: str, api_url: str | None = None, api_key: str = ""):
"""Initialize with file path."""
self._file_path = file_path
self._api_url = api_url

View file

@ -1,5 +1,4 @@
import logging
from typing import Optional
from core.rag.extractor.extractor_base import BaseExtractor
from core.rag.models.document import Document
@ -15,7 +14,7 @@ class UnstructuredXmlExtractor(BaseExtractor):
file_path: Path to the file to load.
"""
def __init__(self, file_path: str, api_url: Optional[str] = None, api_key: str = ""):
def __init__(self, file_path: str, api_url: str | None = None, api_key: str = ""):
"""Initialize with file path."""
self._file_path = file_path
self._api_url = api_url

View file

@ -1,6 +1,6 @@
from collections.abc import Generator
from datetime import datetime
from typing import Any, Optional
from typing import Any
from core.rag.extractor.watercrawl.client import WaterCrawlAPIClient
@ -9,7 +9,7 @@ class WaterCrawlProvider:
def __init__(self, api_key, base_url: str | None = None):
self.client = WaterCrawlAPIClient(api_key, base_url)
def crawl_url(self, url, options: Optional[dict | Any] = None):
def crawl_url(self, url, options: dict | Any | None = None):
options = options or {}
spider_options = {
"max_depth": 1,

View file

@ -1,7 +1,6 @@
"""Abstract interface for document loader implementations."""
from abc import ABC, abstractmethod
from typing import Optional
from configs import dify_config
from core.model_manager import ModelInstance
@ -31,7 +30,7 @@ class BaseIndexProcessor(ABC):
raise NotImplementedError
@abstractmethod
def clean(self, dataset: Dataset, node_ids: Optional[list[str]], with_keywords: bool = True, **kwargs):
def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs):
raise NotImplementedError
@abstractmethod
@ -52,7 +51,7 @@ class BaseIndexProcessor(ABC):
max_tokens: int,
chunk_overlap: int,
separator: str,
embedding_model_instance: Optional[ModelInstance],
embedding_model_instance: ModelInstance | None,
) -> TextSplitter:
"""
Get the NodeParser object according to the processing rule.

View file

@ -1,7 +1,6 @@
"""Paragraph index processor."""
import uuid
from typing import Optional
from core.rag.cleaner.clean_processor import CleanProcessor
from core.rag.datasource.keyword.keyword_factory import Keyword
@ -85,7 +84,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
else:
keyword.add_texts(documents)
def clean(self, dataset: Dataset, node_ids: Optional[list[str]], with_keywords: bool = True, **kwargs):
def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs):
if dataset.indexing_technique == "high_quality":
vector = Vector(dataset)
if node_ids:

View file

@ -1,7 +1,6 @@
"""Paragraph index processor."""
import uuid
from typing import Optional
from configs import dify_config
from core.model_manager import ModelInstance
@ -109,7 +108,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
]
vector.create(formatted_child_documents)
def clean(self, dataset: Dataset, node_ids: Optional[list[str]], with_keywords: bool = True, **kwargs):
def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs):
# node_ids is segment's node_ids
if dataset.indexing_technique == "high_quality":
delete_child_chunks = kwargs.get("delete_child_chunks") or False
@ -187,7 +186,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
document_node: Document,
rules: Rule,
process_rule_mode: str,
embedding_model_instance: Optional[ModelInstance],
embedding_model_instance: ModelInstance | None,
) -> list[ChildDocument]:
if not rules.subchunk_segmentation:
raise ValueError("No subchunk segmentation found in rules.")

View file

@ -4,7 +4,6 @@ import logging
import re
import threading
import uuid
from typing import Optional
import pandas as pd
from flask import Flask, current_app
@ -128,7 +127,7 @@ class QAIndexProcessor(BaseIndexProcessor):
vector = Vector(dataset)
vector.create(documents)
def clean(self, dataset: Dataset, node_ids: Optional[list[str]], with_keywords: bool = True, **kwargs):
def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs):
vector = Vector(dataset)
if node_ids:
vector.delete_by_ids(node_ids)

View file

@ -1,6 +1,6 @@
from abc import ABC, abstractmethod
from collections.abc import Sequence
from typing import Any, Optional
from typing import Any
from pydantic import BaseModel, Field
@ -10,7 +10,7 @@ class ChildDocument(BaseModel):
page_content: str
vector: Optional[list[float]] = None
vector: list[float] | None = None
"""Arbitrary metadata about the page content (e.g., source, relationships to other
documents, etc.).
@ -23,16 +23,16 @@ class Document(BaseModel):
page_content: str
vector: Optional[list[float]] = None
vector: list[float] | None = None
"""Arbitrary metadata about the page content (e.g., source, relationships to other
documents, etc.).
"""
metadata: dict = Field(default_factory=dict)
provider: Optional[str] = "dify"
provider: str | None = "dify"
children: Optional[list[ChildDocument]] = None
children: list[ChildDocument] | None = None
class BaseDocumentTransformer(ABC):

View file

@ -1,5 +1,4 @@
from abc import ABC, abstractmethod
from typing import Optional
from core.rag.models.document import Document
@ -10,9 +9,9 @@ class BaseRerankRunner(ABC):
self,
query: str,
documents: list[Document],
score_threshold: Optional[float] = None,
top_n: Optional[int] = None,
user: Optional[str] = None,
score_threshold: float | None = None,
top_n: int | None = None,
user: str | None = None,
) -> list[Document]:
"""
Run rerank model

View file

@ -1,5 +1,3 @@
from typing import Optional
from core.model_manager import ModelInstance
from core.rag.models.document import Document
from core.rag.rerank.rerank_base import BaseRerankRunner
@ -13,9 +11,9 @@ class RerankModelRunner(BaseRerankRunner):
self,
query: str,
documents: list[Document],
score_threshold: Optional[float] = None,
top_n: Optional[int] = None,
user: Optional[str] = None,
score_threshold: float | None = None,
top_n: int | None = None,
user: str | None = None,
) -> list[Document]:
"""
Run rerank model

View file

@ -1,6 +1,5 @@
import math
from collections import Counter
from typing import Optional
import numpy as np
@ -22,9 +21,9 @@ class WeightRerankRunner(BaseRerankRunner):
self,
query: str,
documents: list[Document],
score_threshold: Optional[float] = None,
top_n: Optional[int] = None,
user: Optional[str] = None,
score_threshold: float | None = None,
top_n: int | None = None,
user: str | None = None,
) -> list[Document]:
"""
Run rerank model

View file

@ -4,7 +4,7 @@ import re
import threading
from collections import Counter, defaultdict
from collections.abc import Generator, Mapping
from typing import Any, Optional, Union, cast
from typing import Any, Union, cast
from flask import Flask, current_app
from sqlalchemy import Float, and_, or_, select, text
@ -85,9 +85,9 @@ class DatasetRetrieval:
show_retrieve_source: bool,
hit_callback: DatasetIndexToolCallbackHandler,
message_id: str,
memory: Optional[TokenBufferMemory] = None,
inputs: Optional[Mapping[str, Any]] = None,
) -> Optional[str]:
memory: TokenBufferMemory | None = None,
inputs: Mapping[str, Any] | None = None,
) -> str | None:
"""
Retrieve dataset.
:param app_id: app_id
@ -290,9 +290,9 @@ class DatasetRetrieval:
model_instance: ModelInstance,
model_config: ModelConfigWithCredentialsEntity,
planning_strategy: PlanningStrategy,
message_id: Optional[str] = None,
metadata_filter_document_ids: Optional[dict[str, list[str]]] = None,
metadata_condition: Optional[MetadataCondition] = None,
message_id: str | None = None,
metadata_filter_document_ids: dict[str, list[str]] | None = None,
metadata_condition: MetadataCondition | None = None,
):
tools = []
for dataset in available_datasets:
@ -410,12 +410,12 @@ class DatasetRetrieval:
top_k: int,
score_threshold: float,
reranking_mode: str,
reranking_model: Optional[dict] = None,
weights: Optional[dict[str, Any]] = None,
reranking_model: dict | None = None,
weights: dict[str, Any] | None = None,
reranking_enable: bool = True,
message_id: Optional[str] = None,
metadata_filter_document_ids: Optional[dict[str, list[str]]] = None,
metadata_condition: Optional[MetadataCondition] = None,
message_id: str | None = None,
metadata_filter_document_ids: dict[str, list[str]] | None = None,
metadata_condition: MetadataCondition | None = None,
):
if not available_datasets:
return []
@ -505,9 +505,7 @@ class DatasetRetrieval:
return all_documents
def _on_retrieval_end(
self, documents: list[Document], message_id: Optional[str] = None, timer: Optional[dict] = None
):
def _on_retrieval_end(self, documents: list[Document], message_id: str | None = None, timer: dict | None = None):
"""Handle retrieval end."""
dify_documents = [document for document in documents if document.provider == "dify"]
for document in dify_documents:
@ -588,8 +586,8 @@ class DatasetRetrieval:
query: str,
top_k: int,
all_documents: list,
document_ids_filter: Optional[list[str]] = None,
metadata_condition: Optional[MetadataCondition] = None,
document_ids_filter: list[str] | None = None,
metadata_condition: MetadataCondition | None = None,
):
with flask_app.app_context():
dataset_stmt = select(Dataset).where(Dataset.id == dataset_id)
@ -664,7 +662,7 @@ class DatasetRetrieval:
hit_callback: DatasetIndexToolCallbackHandler,
user_id: str,
inputs: dict,
) -> Optional[list[DatasetRetrieverBaseTool]]:
) -> list[DatasetRetrieverBaseTool] | None:
"""
A dataset tool is a tool that can be used to retrieve information from a dataset
:param tenant_id: tenant id
@ -853,9 +851,9 @@ class DatasetRetrieval:
user_id: str,
metadata_filtering_mode: str,
metadata_model_config: ModelConfig,
metadata_filtering_conditions: Optional[MetadataFilteringCondition],
metadata_filtering_conditions: MetadataFilteringCondition | None,
inputs: dict,
) -> tuple[Optional[dict[str, list[str]]], Optional[MetadataCondition]]:
) -> tuple[dict[str, list[str]] | None, MetadataCondition | None]:
document_query = db.session.query(DatasetDocument).where(
DatasetDocument.dataset_id.in_(dataset_ids),
DatasetDocument.indexing_status == "completed",
@ -950,7 +948,7 @@ class DatasetRetrieval:
def _automatic_metadata_filter_func(
self, dataset_ids: list, query: str, tenant_id: str, user_id: str, metadata_model_config: ModelConfig
) -> Optional[list[dict[str, Any]]]:
) -> list[dict[str, Any]] | None:
# get all metadata field
metadata_stmt = select(DatasetMetadata).where(DatasetMetadata.dataset_id.in_(dataset_ids))
metadata_fields = db.session.scalars(metadata_stmt).all()
@ -1005,7 +1003,7 @@ class DatasetRetrieval:
return automatic_metadata_filters
def _process_metadata_filter_func(
self, sequence: int, condition: str, metadata_name: str, value: Optional[Any], filters: list
self, sequence: int, condition: str, metadata_name: str, value: Any | None, filters: list
):
if value is None and condition not in ("empty", "not empty"):
return

View file

@ -2,7 +2,7 @@
from __future__ import annotations
from typing import Any, Optional
from typing import Any
from core.model_manager import ModelInstance
from core.model_runtime.model_providers.__base.tokenizers.gpt2_tokenizer import GPT2Tokenizer
@ -24,7 +24,7 @@ class EnhanceRecursiveCharacterTextSplitter(RecursiveCharacterTextSplitter):
@classmethod
def from_encoder(
cls: type[TS],
embedding_model_instance: Optional[ModelInstance],
embedding_model_instance: ModelInstance | None,
allowed_special: Union[Literal["all"], Set[str]] = set(), # noqa: UP037
disallowed_special: Union[Literal["all"], Collection[str]] = "all", # noqa: UP037
**kwargs: Any,
@ -48,7 +48,7 @@ class EnhanceRecursiveCharacterTextSplitter(RecursiveCharacterTextSplitter):
class FixedRecursiveCharacterTextSplitter(EnhanceRecursiveCharacterTextSplitter):
def __init__(self, fixed_separator: str = "\n\n", separators: Optional[list[str]] = None, **kwargs: Any):
def __init__(self, fixed_separator: str = "\n\n", separators: list[str] | None = None, **kwargs: Any):
"""Create a new TextSplitter."""
super().__init__(**kwargs)
self._fixed_separator = fixed_separator

View file

@ -9,7 +9,6 @@ from dataclasses import dataclass
from typing import (
Any,
Literal,
Optional,
TypeVar,
Union,
)
@ -71,7 +70,7 @@ class TextSplitter(BaseDocumentTransformer, ABC):
def split_text(self, text: str) -> list[str]:
"""Split text into multiple components."""
def create_documents(self, texts: list[str], metadatas: Optional[list[dict]] = None) -> list[Document]:
def create_documents(self, texts: list[str], metadatas: list[dict] | None = None) -> list[Document]:
"""Create documents from a list of texts."""
_metadatas = metadatas or [{}] * len(texts)
documents = []
@ -94,7 +93,7 @@ class TextSplitter(BaseDocumentTransformer, ABC):
metadatas.append(doc.metadata or {})
return self.create_documents(texts, metadatas=metadatas)
def _join_docs(self, docs: list[str], separator: str) -> Optional[str]:
def _join_docs(self, docs: list[str], separator: str) -> str | None:
text = separator.join(docs)
text = text.strip()
if text == "":
@ -194,7 +193,7 @@ class TokenTextSplitter(TextSplitter):
def __init__(
self,
encoding_name: str = "gpt2",
model_name: Optional[str] = None,
model_name: str | None = None,
allowed_special: Union[Literal["all"], Set[str]] = set(),
disallowed_special: Union[Literal["all"], Collection[str]] = "all",
**kwargs: Any,
@ -245,7 +244,7 @@ class RecursiveCharacterTextSplitter(TextSplitter):
def __init__(
self,
separators: Optional[list[str]] = None,
separators: list[str] | None = None,
keep_separator: bool = True,
**kwargs: Any,
):