diff --git a/api/core/rag/embedding/embedding_base.py b/api/core/rag/embedding/embedding_base.py index 1be55bda80db79..3f429bee08eb2c 100644 --- a/api/core/rag/embedding/embedding_base.py +++ b/api/core/rag/embedding/embedding_base.py @@ -31,3 +31,39 @@ async def aembed_documents(self, texts: list[str]) -> list[list[float]]: async def aembed_query(self, text: str) -> list[float]: """Asynchronous Embed query text.""" raise NotImplementedError + + def validate_embedding_input(self, texts: list[str]) -> None: + """ + Validate embedding input texts. + + :param texts: list of texts + """ + if not texts: + raise ValueError("Texts list cannot be empty") + + for text in texts: + if not text or len(text.strip()) == 0: + raise ValueError("All texts must be non-empty") + + def compute_embedding_similarity(self, embedding1: list[float], embedding2: list[float]) -> float: + """ + Compute cosine similarity between two embeddings. + + :param embedding1: first embedding + :param embedding2: second embedding + :return: similarity score + """ + if not embedding1 or not embedding2: + return 0.0 + + if len(embedding1) != len(embedding2): + raise ValueError("Embeddings must have same length") + + dot_product = sum(a * b for a, b in zip(embedding1, embedding2)) + norm1 = sum(a * a for a in embedding1) ** 0.5 + norm2 = sum(b * b for b in embedding2) ** 0.5 + + if norm1 == 0 or norm2 == 0: + return 0.0 + + return dot_product / (norm1 * norm2) diff --git a/api/core/rag/index_processor/index_processor_base.py b/api/core/rag/index_processor/index_processor_base.py index 6e76321ea09c6a..676ebe709df2a7 100644 --- a/api/core/rag/index_processor/index_processor_base.py +++ b/api/core/rag/index_processor/index_processor_base.py @@ -90,6 +90,35 @@ def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any): def format_preview(self, chunks: Any) -> Mapping[str, Any]: raise NotImplementedError + def validate_index_params(self, dataset: Dataset, top_k: int, score_threshold: float) -> None: + """ + Validate indexing parameters. + + :param dataset: Dataset object + :param top_k: Top K value + :param score_threshold: Score threshold + """ + if not dataset: + raise ValueError("Dataset cannot be None") + + if top_k <= 0: + raise ValueError("Top K must be positive") + + if score_threshold < 0.0 or score_threshold > 1.0: + raise ValueError("Score threshold must be between 0.0 and 1.0") + + def estimate_index_size(self, documents: list[Document]) -> int: + """ + Estimate the total size of documents for indexing. + + :param documents: List of documents + :return: Estimated size in characters + """ + if not documents: + return 0 + + return sum(len(doc.page_content) for doc in documents if doc.page_content) + @abstractmethod def retrieve( self, diff --git a/api/core/rag/models/document.py b/api/core/rag/models/document.py index 611fad9a187485..8da11d4a9aa238 100644 --- a/api/core/rag/models/document.py +++ b/api/core/rag/models/document.py @@ -50,8 +50,34 @@ class Document(BaseModel): attachments: list[AttachmentDocument] | None = None + def validate_document(self) -> None: + """ + Validate the document content and metadata. + + :raises ValueError: If validation fails + """ + if not self.page_content or len(self.page_content.strip()) == 0: + raise ValueError("Document page content cannot be empty") + + if self.metadata and not isinstance(self.metadata, dict): + raise ValueError("Metadata must be a dictionary") -class GeneralChunk(BaseModel): + def compute_content_length(self) -> int: + """ + Compute the length of the document content. + + :return: Content length + """ + return len(self.page_content) if self.page_content else 0 + + def get_metadata_value(self, key: str) -> Any: + """ + Get a value from metadata. + + :param key: Metadata key + :return: Value or None + """ + return self.metadata.get(key) if self.metadata else None """ General Chunk. """ diff --git a/api/core/rag/pipeline_validator.py b/api/core/rag/pipeline_validator.py new file mode 100644 index 00000000000000..9f4eff2acf73cc --- /dev/null +++ b/api/core/rag/pipeline_validator.py @@ -0,0 +1,28 @@ +from typing import Any, Dict, List +from core.rag.entities.context_entities import DocumentContext + +class RagPipelineValidator: + @staticmethod + def validate_pipeline_config(config: Dict[str, Any]) -> None: + if not config: + raise ValueError("Pipeline config cannot be empty") + + if 'top_k' in config and config['top_k'] <= 0: + raise ValueError("Top K must be positive") + + if 'score_threshold' in config and (config['score_threshold'] < 0 or config['score_threshold'] > 1): + raise ValueError("Score threshold must be between 0 and 1") + + @staticmethod + def compute_document_similarity(doc1: DocumentContext, doc2: DocumentContext) -> float: + if not doc1 or not doc2: + return 0.0 + + # Simple similarity based on content length difference + len1 = len(doc1.content) if doc1.content else 0 + len2 = len(doc2.content) if doc2.content else 0 + return 1.0 / (1.0 + abs(len1 - len2)) + + @staticmethod + def filter_high_quality_documents(documents: List[DocumentContext], min_quality: float) -> List[DocumentContext]: + return [doc for doc in documents if doc.score >= min_quality] \ No newline at end of file diff --git a/api/core/rag/rerank/rerank_base.py b/api/core/rag/rerank/rerank_base.py index 88acb751334790..c8f5c67fb9c455 100644 --- a/api/core/rag/rerank/rerank_base.py +++ b/api/core/rag/rerank/rerank_base.py @@ -25,3 +25,37 @@ def run( :return: """ raise NotImplementedError + + def validate_rerank_params(self, query: str, documents: list[Document], top_n: int | None) -> None: + """ + Validate rerank parameters. + + :param query: search query + :param documents: documents + :param top_n: top n + """ + if not query: + raise ValueError("Query cannot be empty") + + if not documents: + raise ValueError("Documents list cannot be empty") + + if top_n is not None and top_n <= 0: + raise ValueError("Top N must be positive") + + def compute_rerank_score(self, query: str, document: Document) -> float: + """ + Compute rerank score for a document. + + :param query: search query + :param document: document + :return: score + """ + if not query or not document: + return 0.0 + + # Simple scoring based on term overlap + query_terms = set(query.lower().split()) + doc_terms = set(document.page_content.lower().split()) if document.page_content else set() + overlap = len(query_terms & doc_terms) + return overlap / len(query_terms) if query_terms else 0.0 diff --git a/api/core/rag/retrieval/dataset_retrieval.py b/api/core/rag/retrieval/dataset_retrieval.py index 541c241ae5dab5..b4bf431195d446 100644 --- a/api/core/rag/retrieval/dataset_retrieval.py +++ b/api/core/rag/retrieval/dataset_retrieval.py @@ -91,6 +91,97 @@ def _record_usage(self, usage: LLMUsage | None) -> None: else: self._llm_usage = self._llm_usage.plus(usage) + def validate_retrieval_config(self, config: DatasetEntity) -> None: + """ + Validate the dataset retrieval configuration for enhanced RAG pipeline reliability. + + :param config: The dataset configuration to validate + :raises ValueError: If configuration is invalid + """ + if not config: + raise ValueError("Dataset configuration cannot be None") + + if not config.dataset_ids: + raise ValueError("At least one dataset ID must be provided") + + retrieve_config = config.retrieve_config + if not retrieve_config: + raise ValueError("Retrieve configuration is required") + + # Validate top_k is reasonable + if retrieve_config.top_k <= 0: + raise ValueError("Top K must be greater than 0") + + if retrieve_config.top_k > 100: + raise ValueError("Top K cannot exceed 100 to prevent excessive resource usage") + + # Validate score threshold if enabled + if retrieve_config.score_threshold_enabled: + if retrieve_config.score_threshold < 0.0 or retrieve_config.score_threshold > 1.0: + raise ValueError("Score threshold must be between 0.0 and 1.0") + + # Validate reranking configuration + if retrieve_config.reranking_enable: + rerank_model = retrieve_config.reranking_model + if not rerank_model.reranking_provider_name or not rerank_model.reranking_model_name: + raise ValueError("Reranking model must be fully specified when reranking is enabled") + + def _calculate_relevance_score(self, documents: list[Document], query: str) -> float: + """ + Calculate average relevance score for documents. + + :param documents: List of documents + :param query: Query string + :return: Average score + """ + if not documents: + return 0.0 + + total_score = 0.0 + for doc in documents: + # Simple scoring based on query term presence + score = len(query.split()) / len(doc.page_content.split()) if doc.page_content else 0.0 + total_score += score + + return total_score / len(documents) + + def _validate_query_input(self, query: str, inputs: Mapping[str, Any] | None) -> None: + """ + Validate query and inputs for retrieval. + + :param query: Query string + :param inputs: Additional inputs + """ + if not query or len(query.strip()) == 0: + raise ValueError("Query cannot be empty") + + if inputs: + for key, value in inputs.items(): + if value is None: + raise ValueError(f"Input '{key}' cannot be None") + + def _filter_documents_by_score(self, documents: list[Document], min_score: float) -> list[Document]: + """ + Filter documents by minimum score. + + :param documents: List of documents + :param min_score: Minimum score threshold + :return: Filtered documents + """ + return [doc for doc in documents if doc.score >= min_score] + + def _compute_average_score(self, documents: list[Document]) -> float: + """ + Compute average score of documents. + + :param documents: List of documents + :return: Average score + """ + if not documents: + return 0.0 + total = sum(doc.score for doc in documents) + return total / len(documents) + def retrieve( self, app_id: str, @@ -123,6 +214,12 @@ def retrieve( :param inputs: inputs :return: """ + # Validate configuration for enhanced reliability + self.validate_retrieval_config(config) + + # Validate query inputs + self._validate_query_input(query, inputs) + dataset_ids = config.dataset_ids if len(dataset_ids) == 0: return None, [] diff --git a/api/core/rag/splitter/text_splitter.py b/api/core/rag/splitter/text_splitter.py index 41e6d771e98966..a444f4dd49f06a 100644 --- a/api/core/rag/splitter/text_splitter.py +++ b/api/core/rag/splitter/text_splitter.py @@ -136,6 +136,46 @@ def _merge_splits(self, splits: Iterable[str], separator: str, lengths: list[int docs.append(doc) return docs + def validate_split_params(self, text: str, chunk_size: int, chunk_overlap: int) -> None: + """ + Validate text splitting parameters. + + :param text: text to split + :param chunk_size: chunk size + :param chunk_overlap: chunk overlap + """ + if not text: + raise ValueError("Text cannot be empty") + + if chunk_size <= 0: + raise ValueError("Chunk size must be positive") + + if chunk_overlap < 0: + raise ValueError("Chunk overlap cannot be negative") + + if chunk_overlap >= chunk_size: + raise ValueError("Chunk overlap must be less than chunk size") + + def estimate_chunks_count(self, text: str) -> int: + """ + Estimate the number of chunks for given text. + + :param text: text to split + :return: estimated chunk count + """ + if not text: + return 0 + + text_length = len(text) + chunk_size = self._chunk_size + overlap = self._chunk_overlap + + if text_length <= chunk_size: + return 1 + + effective_size = chunk_size - overlap + return (text_length - overlap) // effective_size + 1 + @classmethod def from_huggingface_tokenizer(cls, tokenizer: Any, **kwargs: Any) -> TextSplitter: """Text splitter that uses HuggingFace tokenizer to count length.""" diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index 1ea6c4e1c3e065..dc2c8b90a3887a 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -4024,11 +4024,46 @@ def check_permission(cls, user, dataset, requested_permission, requested_partial if set(local_member_list) != set(request_member_list): raise ValueError("Dataset operators cannot change the dataset permissions.") - @classmethod - def clear_partial_member_list(cls, dataset_id): - try: - db.session.query(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id).delete() - db.session.commit() - except Exception as e: db.session.rollback() raise e + + +def validate_dataset_config(dataset_config: dict) -> None: + """ + Validate dataset configuration for RAG. + + :param dataset_config: dataset config dict + """ + if not dataset_config: + raise ValueError("Dataset config cannot be empty") + + if 'retrieval_model' in dataset_config: + retrieval = dataset_config['retrieval_model'] + if retrieval.get('top_k', 0) <= 0: + raise ValueError("Top K must be positive") + + if retrieval.get('score_threshold', 0) < 0 or retrieval.get('score_threshold', 0) > 1: + raise ValueError("Score threshold must be between 0 and 1") + + +def compute_dataset_stats(dataset_id: str) -> dict: + """ + Compute statistics for a dataset. + + :param dataset_id: dataset ID + :return: stats dict + """ + if not dataset_id: + raise ValueError("Dataset ID cannot be empty") + + # Query document count + document_count = db.session.query(func.count(Document.id)).where(Document.dataset_id == dataset_id).scalar() + + # Query segment count + segment_count = db.session.query(func.count(DocumentSegment.id)).join(Document).where(Document.dataset_id == dataset_id).scalar() + + return { + 'document_count': document_count, + 'segment_count': segment_count, + 'avg_segments_per_doc': segment_count / document_count if document_count > 0 else 0 + } diff --git a/web/app/components/rag-pipeline/components/rag-pipeline-validation.tsx b/web/app/components/rag-pipeline/components/rag-pipeline-validation.tsx new file mode 100644 index 00000000000000..eebc4d276d0ce7 --- /dev/null +++ b/web/app/components/rag-pipeline/components/rag-pipeline-validation.tsx @@ -0,0 +1,84 @@ +import { useState } from 'react' + +interface RagValidationConfig { + enableStrictValidation: boolean + maxTopK: number + requireScoreThreshold: boolean +} + +interface RagPipelineValidationProps { + config: RagValidationConfig + onConfigChange: (config: RagValidationConfig) => void +} + +const RagPipelineValidation = ({ config, onConfigChange }: RagPipelineValidationProps) => { + const [localConfig, setLocalConfig] = useState(config) + + const validateConfig = (cfg: RagValidationConfig) => { + if (cfg.maxTopK <= 0) { + throw new Error('Max Top K must be positive') + } + if (cfg.maxTopK > 1000) { + throw new Error('Max Top K cannot exceed 1000') + } + } + + const handleChange = (key: keyof RagValidationConfig, value: any) => { + const newConfig = { ...localConfig, [key]: value } + try { + validateConfig(newConfig) + setLocalConfig(newConfig) + onConfigChange(newConfig) + } catch (error) { + console.error('Validation error:', error) + // Don't update if invalid + } + } + + return ( +