Test runner for query retrieval tests.
| 52 | |
| 53 | |
| 54 | class QueryRetrievalTestRunner(BaseTestRunner): |
| 55 | """Test runner for query retrieval tests.""" |
| 56 | |
| 57 | def __init__(self): |
| 58 | """Initialize query retrieval test runner.""" |
| 59 | super().__init__(TestType.QUERY_RETRIEVAL) |
| 60 | |
| 61 | def validate_test_case(self, test_case: dict[str, Any]) -> tuple[bool, str | None]: |
| 62 | """Validate query retrieval test case has required fields.""" |
| 63 | required_fields = ['query', 'retrieval_backend'] |
| 64 | missing_fields = [field for field in required_fields if field not in test_case or not test_case[field]] |
| 65 | |
| 66 | if missing_fields: |
| 67 | return False, f"Missing required fields: {missing_fields}" |
| 68 | |
| 69 | # Validate expected_urls format if present |
| 70 | if test_case.get('expected_urls') and not isinstance(test_case['expected_urls'], list): |
| 71 | return False, "expected_urls must be a list" |
| 72 | |
| 73 | # Validate contains_urls format if present |
| 74 | if test_case.get('contains_urls') and not isinstance(test_case['contains_urls'], list): |
| 75 | return False, "contains_urls must be a list" |
| 76 | |
| 77 | # Validate excludes_urls format if present |
| 78 | if test_case.get('excludes_urls') and not isinstance(test_case['excludes_urls'], list): |
| 79 | return False, "excludes_urls must be a list" |
| 80 | |
| 81 | return True, None |
| 82 | |
| 83 | async def run_single_test(self, test_case: dict[str, Any]) -> QueryRetrievalTestResult: |
| 84 | """Run a single query retrieval test.""" |
| 85 | start_time = time.time() |
| 86 | |
| 87 | # Create test case object |
| 88 | query_case = QueryRetrievalTestCase( |
| 89 | test_type=self.test_type, |
| 90 | test_id=test_case.get('test_id', 0), |
| 91 | original_case_num=test_case.get('original_case_num', 0), |
| 92 | query=test_case['query'], |
| 93 | retrieval_backend=test_case.get('db', test_case.get('retrieval_backend')), |
| 94 | site=test_case.get('site', 'all'), |
| 95 | num_results=test_case.get('num_results', test_case.get('top_k', 10)), # Support both for backwards compatibility |
| 96 | expected_min_results=test_case.get('expected_min_results'), |
| 97 | expected_max_results=test_case.get('expected_max_results'), |
| 98 | expected_urls=test_case.get('expected_urls'), |
| 99 | contains_urls=test_case.get('contains_urls'), |
| 100 | excludes_urls=test_case.get('excludes_urls'), |
| 101 | min_score=test_case.get('min_score'), |
| 102 | description=test_case.get('description', f"Query: {test_case['query'][:50]}...") |
| 103 | ) |
| 104 | |
| 105 | try: |
| 106 | logger.info(f"Starting query retrieval test for: '{query_case.query}'") |
| 107 | logger.debug(f"Parameters: backend={query_case.retrieval_backend}, site={query_case.site}, num_results={query_case.num_results}") |
| 108 | |
| 109 | # Create VectorDBClient instance |
| 110 | client = VectorDBClient(endpoint_name=query_case.retrieval_backend) |
| 111 |