Test runner for site retrieval tests.
| 47 | |
| 48 | |
| 49 | class SiteRetrievalTestRunner(BaseTestRunner): |
| 50 | """Test runner for site retrieval tests.""" |
| 51 | |
| 52 | def __init__(self): |
| 53 | """Initialize site retrieval test runner.""" |
| 54 | super().__init__(TestType.SITE_RETRIEVAL) |
| 55 | |
| 56 | def validate_test_case(self, test_case: dict[str, Any]) -> tuple[bool, str | None]: |
| 57 | """Validate site retrieval test case has required fields.""" |
| 58 | required_fields = ['retrieval_backend'] |
| 59 | missing_fields = [field for field in required_fields if field not in test_case or not test_case[field]] |
| 60 | |
| 61 | if missing_fields: |
| 62 | return False, f"Missing required fields: {missing_fields}" |
| 63 | |
| 64 | # Validate expected_sites format if present |
| 65 | if test_case.get('expected_sites') and not isinstance(test_case['expected_sites'], list): |
| 66 | return False, "expected_sites must be a list" |
| 67 | |
| 68 | # Validate contains_sites format if present |
| 69 | if test_case.get('contains_sites') and not isinstance(test_case['contains_sites'], list): |
| 70 | return False, "contains_sites must be a list" |
| 71 | |
| 72 | # Validate excludes_sites format if present |
| 73 | if test_case.get('excludes_sites') and not isinstance(test_case['excludes_sites'], list): |
| 74 | return False, "excludes_sites must be a list" |
| 75 | |
| 76 | return True, None |
| 77 | |
| 78 | async def run_single_test(self, test_case: dict[str, Any]) -> SiteRetrievalTestResult: |
| 79 | """Run a single site retrieval test.""" |
| 80 | start_time = time.time() |
| 81 | |
| 82 | # Create test case object |
| 83 | site_case = SiteRetrievalTestCase( |
| 84 | test_type=self.test_type, |
| 85 | test_id=test_case.get('test_id', 0), |
| 86 | original_case_num=test_case.get('original_case_num', 0), |
| 87 | retrieval_backend=test_case.get('db', test_case.get('retrieval_backend')), |
| 88 | expected_sites=test_case.get('expected_sites'), |
| 89 | expected_min_sites=test_case.get('expected_min_sites'), |
| 90 | expected_max_sites=test_case.get('expected_max_sites'), |
| 91 | contains_sites=test_case.get('contains_sites'), |
| 92 | excludes_sites=test_case.get('excludes_sites'), |
| 93 | description=test_case.get('description', f"Site retrieval test for {test_case.get('retrieval_backend')}") |
| 94 | ) |
| 95 | |
| 96 | try: |
| 97 | logger.info(f"Starting site retrieval test for backend: {site_case.retrieval_backend}") |
| 98 | |
| 99 | # Create VectorDBClient instance |
| 100 | client = VectorDBClient(endpoint_name=site_case.retrieval_backend) |
| 101 | |
| 102 | # Get sites |
| 103 | sites = await client.get_sites() |
| 104 | site_count = len(sites) |
| 105 | |
| 106 | logger.info(f"Retrieved {site_count} sites from {site_case.retrieval_backend}") |