-
Notifications
You must be signed in to change notification settings - Fork 2
Feature/enhanced rag validation #10
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
12c8639
66d31b0
79fceb0
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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] | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -25,3 +25,37 @@ def run( | |
| :return: | ||
| """ | ||
| raise NotImplementedError | ||
|
|
||
| def validate_rerank_params(self, query: str, documents: list[Document], top_n: int | None) -> None: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| """ | ||
| 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: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| """ | ||
| 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 | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| 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: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| """ | ||
| 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 | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| total_score += score | ||
|
|
||
| return total_score / len(documents) | ||
|
|
||
| def _validate_query_input(self, query: str, inputs: Mapping[str, Any] | None) -> None: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| """ | ||
| 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]: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| """ | ||
| 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: | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| """ | ||
| 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, [] | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
|
Comment on lines
+176
to
+177
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
|
|
||
| @classmethod | ||
| def from_huggingface_tokenizer(cls, tokenizer: Any, **kwargs: Any) -> TextSplitter: | ||
| """Text splitter that uses HuggingFace tokenizer to count length.""" | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -4024,11 +4024,46 @@ | |
| 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 | ||
|
Check failure on line 4028 in api/services/dataset_service.py
|
||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
|
|
||
|
|
||
| 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 | ||
| } | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The
compute_document_similarityfunction bases its calculation only on the length of document content, which can be misleading. For example, two completely different documents are rated as identical if they have the same character count, potentially causing incorrect behavior in features relying on this score.Rename the function to be more descriptive, like
compute_length_based_similarity, or implement a content-based similarity algorithm such as Jaccard similarity to accurately reflect document content similarity.