File size: 16,359 Bytes
18a82fb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445

# !/usr/bin/env python3
# Add this comprehensive test function to your url_dataset.py file
# Replace the existing comprehensive test with this fixed version

def run_comprehensive_test():
    """Run comprehensive test suite for url_dataset.py"""
    print("URL_DATASET.PY COMPREHENSIVE TEST SUITE")
    print("=" * 60)
    print("Testing all classes and functions in url_dataset.py")

    import tempfile
    import json
    import hashlib
    from unittest.mock import patch, MagicMock
    from io import BytesIO
    import requests
    from collections import Counter
    from torchvision import transforms

    tests_passed = 0
    tests_failed = 0
    failed_tests = []

    def run_test(test_name, test_func):
        nonlocal tests_passed, tests_failed, failed_tests
        print(f"\n{'=' * 60}")
        print(f"Running: {test_name}")
        print(f"{'=' * 60}")
        try:
            test_func()
            tests_passed += 1
            print(f"βœ… PASSED: {test_name}")
        except Exception as e:
            tests_failed += 1
            failed_tests.append(f"{test_name}: {str(e)}")
            print(f"❌ FAILED: {test_name}")
            print(f"Error: {str(e)}")

    def create_mock_data(num_items=10):
        """Create mock JSON data"""
        mock_data = []
        decades = ['1960s', '1970s', '1980s', '1990s', '2000s']
        for i in range(num_items):
            mock_data.append({
                "id": f"test_id_{i}",
                "product_id": f"product_{i % 3}",
                "name": f"Test Product {i}",
                "decade": decades[i % 5],
                "url": f"https://example.com/image_{i}.jpg",
                "classification": f"test_class_{i % 2}",
                "makers": f"test_maker_{i % 2}",
                "country": f"test_country_{i % 3}"
            })
        return mock_data

    def create_mock_image(size=(200, 200)):
        """Create mock PIL image"""
        array = np.random.randint(0, 255, (*size, 3), dtype=np.uint8)
        return Image.fromarray(array)

    def test_base_dataset():
        """Test BaseDataset functionality"""
        print("Testing BaseDataset...")

        with tempfile.TemporaryDirectory() as temp_dir:
            temp_path = Path(temp_dir)
            split_file = temp_path / "test.json"

            mock_data = create_mock_data(10)
            with open(split_file, 'w') as f:
                json.dump(mock_data, f)

            dataset = BaseDataset(str(split_file))

            print(f"βœ“ BaseDataset initialized with {len(dataset)} items")
            print(f"βœ“ Number of classes: {dataset.num_classes}")
            print(f"βœ“ Class names: {dataset.decades}")

            assert len(dataset) == 10
            assert dataset.num_classes == 5
            assert dataset.decades == ['1960s', '1970s', '1980s', '1990s', '2000s']

            labels = dataset.get_labels()
            print(f"βœ“ Labels: {labels}")
            assert len(labels) == 10

            metadata = dataset.get_metadata(0)
            expected_keys = ['id', 'product_id', 'name', 'decade', 'url', 'classification', 'makers', 'country']
            for key in expected_keys:
                assert key in metadata
            print(f"βœ“ Metadata keys: {list(metadata.keys())}")

    def test_url_dataset_mock():
        """Test URLDataset with mocked requests"""
        print("Testing URLDataset with mocked requests...")

        with tempfile.TemporaryDirectory() as temp_dir:
            temp_path = Path(temp_dir)
            cache_dir = temp_path / "cache"
            split_file = temp_path / "test.json"

            mock_data = create_mock_data(5)
            with open(split_file, 'w') as f:
                json.dump(mock_data, f)

            # Create transform to ensure tensor output
            test_transform = transforms.Compose([
                transforms.Resize((224, 224)),
                transforms.ToTensor()
            ])

            mock_image = create_mock_image((300, 300))
            mock_response = MagicMock()
            mock_response.content = BytesIO()
            mock_image.save(mock_response.content, 'JPEG')
            mock_response.content = mock_response.content.getvalue()
            mock_response.raise_for_status = MagicMock()

            with patch('requests.get', return_value=mock_response):
                dataset = URLDataset(
                    split_file=str(split_file),
                    cache_dir=str(cache_dir),
                    transform=test_transform,  # Provide transform
                    max_retries=2,
                    timeout=5,
                    fallback_on_error=True
                )

                print(f"βœ“ URLDataset initialized with {len(dataset)} items")
                print(f"βœ“ Cache directory: {dataset.cache_dir}")

                test_url = "https://example.com/test.jpg"
                cache_path = dataset._get_cache_path(test_url)
                print(f"βœ“ Cache path generation: {cache_path.name}")

                # Test image loading
                image, label, metadata = dataset[0]

                assert isinstance(image, torch.Tensor)
                assert image.dtype == torch.float32
                assert 0 <= label < 5
                assert isinstance(metadata, dict)

                print(f"βœ“ First item loaded: shape={image.shape}, label={label}, decade={metadata['decade']}")

                # Test statistics
                stats = dataset.get_statistics()
                print(f"βœ“ Statistics: {stats}")

    def test_url_dataset_fallback():
        """Test URLDataset fallback behavior"""
        print("Testing URLDataset fallback behavior...")

        with tempfile.TemporaryDirectory() as temp_dir:
            temp_path = Path(temp_dir)
            cache_dir = temp_path / "cache"
            split_file = temp_path / "test.json"

            mock_data = create_mock_data(3)
            with open(split_file, 'w') as f:
                json.dump(mock_data, f)

            # Create transform
            test_transform = transforms.Compose([
                transforms.Resize((224, 224)),
                transforms.ToTensor()
            ])

            with patch('requests.get', side_effect=requests.exceptions.ConnectionError("Mock connection error")):
                dataset = URLDataset(
                    split_file=str(split_file),
                    cache_dir=str(cache_dir),
                    transform=test_transform,  # Provide transform
                    max_retries=1,
                    timeout=1,
                    fallback_on_error=True
                )

                # Should use placeholder image
                image, label, metadata = dataset[0]

                assert isinstance(image, torch.Tensor)
                print(f"βœ“ Fallback image loaded: shape={image.shape}")

                stats = dataset.get_statistics()
                assert stats['failures'] > 0
                print(f"βœ“ Failure tracked in statistics: {stats['failures']} failures")

    def test_cached_dataset():
        """Test CachedDataset functionality"""
        print("Testing CachedDataset...")

        with tempfile.TemporaryDirectory() as temp_dir:
            temp_path = Path(temp_dir)
            images_dir = temp_path / "images"
            images_dir.mkdir()
            split_file = temp_path / "test.json"

            mock_data = create_mock_data(5)
            with open(split_file, 'w') as f:
                json.dump(mock_data, f)

            # Create cached images for first 3 items
            for i in range(3):
                url = mock_data[i]['url']
                url_hash = hashlib.md5(url.encode()).hexdigest()
                cache_path = images_dir / f"{url_hash}.jpg"

                mock_image = create_mock_image((200, 200))
                mock_image.save(cache_path, 'JPEG')

            print(f"βœ“ Created {len(list(images_dir.glob('*.jpg')))} cached images")

            # Create transform
            test_transform = transforms.Compose([
                transforms.Resize((224, 224)),
                transforms.ToTensor()
            ])

            dataset = CachedDataset(
                split_file=str(split_file),
                images_dir=str(images_dir),
                transform=test_transform,  # Provide transform
                verify_images=True
            )

            print(f"βœ“ CachedDataset initialized with {len(dataset)} valid images")
            assert len(dataset) == 3

            image, label, metadata = dataset[0]
            assert isinstance(image, torch.Tensor)
            assert 0 <= label < 5
            print(f"βœ“ Cached item loaded: shape={image.shape}, label={label}")

    def test_subset_creation():
        """Test subset creation"""
        print("Testing create_subset_dataset...")

        with tempfile.TemporaryDirectory() as temp_dir:
            temp_path = Path(temp_dir)
            split_file = temp_path / "test.json"

            mock_data = create_mock_data(50)
            with open(split_file, 'w') as f:
                json.dump(mock_data, f)

            original_dataset = BaseDataset(str(split_file))
            original_size = len(original_dataset)

            print(f"βœ“ Original dataset size: {original_size}")

            subset_dataset = create_subset_dataset(original_dataset, fraction=0.2, seed=42)
            subset_size = len(subset_dataset)

            print(f"βœ“ Subset dataset size: {subset_size}")
            print(f"βœ“ Subset fraction: {subset_size / original_size:.2f}")

            assert subset_size >= 5  # At least 1 from each class
            assert subset_size <= original_size

            subset_labels = subset_dataset.get_labels()
            unique_labels = set(subset_labels)
            print(f"βœ“ Subset has {len(unique_labels)} unique classes: {unique_labels}")

    def test_download_images():
        """Test bulk download functionality"""
        print("Testing download_dataset_images with mocked requests...")

        with tempfile.TemporaryDirectory() as temp_dir:
            temp_path = Path(temp_dir)
            cache_dir = temp_path / "cache"

            split_files = []
            for split_name in ["train", "val"]:
                split_file = temp_path / f"{split_name}.json"
                mock_data = create_mock_data(5)
                # Make URLs unique across splits
                for i, item in enumerate(mock_data):
                    item['url'] = f"https://example.com/{split_name}_image_{i}.jpg"

                with open(split_file, 'w') as f:
                    json.dump(mock_data, f)
                split_files.append(str(split_file))

            mock_image = create_mock_image((200, 200))
            mock_response = MagicMock()
            mock_response.content = BytesIO()
            mock_image.save(mock_response.content, 'JPEG')
            mock_response.content = mock_response.content.getvalue()
            mock_response.raise_for_status = MagicMock()

            with patch('requests.get', return_value=mock_response):
                results = download_dataset_images(
                    split_files=split_files,
                    cache_dir=str(cache_dir),
                    num_workers=2,
                    skip_existing=True
                )

                print(f"βœ“ Download results: {results}")

                expected_keys = ['cached', 'downloaded', 'failed']
                for key in expected_keys:
                    assert key in results

                cached_images = list(cache_dir.glob('*.jpg'))
                print(f"βœ“ Created {len(cached_images)} cached image files")

    def test_edge_cases():
        """Test edge cases"""
        print("Testing edge cases...")

        with tempfile.TemporaryDirectory() as temp_dir:
            temp_path = Path(temp_dir)

            # Test empty dataset
            empty_split_file = temp_path / "empty.json"
            with open(empty_split_file, 'w') as f:
                json.dump([], f)

            empty_dataset = BaseDataset(str(empty_split_file))
            assert len(empty_dataset) == 0
            print("βœ“ Empty dataset handled correctly")

            # Test malformed JSON
            try:
                malformed_split_file = temp_path / "malformed.json"
                with open(malformed_split_file, 'w') as f:
                    f.write("invalid json")

                BaseDataset(str(malformed_split_file))
                assert False, "Should have raised exception"
            except json.JSONDecodeError:
                print("βœ“ Malformed JSON handled correctly")

            # Test URLDataset with no fallback
            split_file = temp_path / "test.json"
            mock_data = create_mock_data(2)
            with open(split_file, 'w') as f:
                json.dump(mock_data, f)

            test_transform = transforms.Compose([
                transforms.Resize((224, 224)),
                transforms.ToTensor()
            ])

            with patch('requests.get', side_effect=Exception("Mock error")):
                dataset = URLDataset(
                    split_file=str(split_file),
                    transform=test_transform,
                    fallback_on_error=False,
                    max_retries=1
                )

                try:
                    image, label, metadata = dataset[0]
                    assert False, "Should have raised exception"
                except (ValueError, Exception):
                    print("βœ“ No-fallback error handling works correctly")

    def test_performance():
        """Test performance"""
        print("Testing performance...")

        with tempfile.TemporaryDirectory() as temp_dir:
            temp_path = Path(temp_dir)
            split_file = temp_path / "perf_test.json"

            mock_data = create_mock_data(100)
            with open(split_file, 'w') as f:
                json.dump(mock_data, f)

            import time

            start_time = time.time()
            dataset = BaseDataset(str(split_file))
            init_time = time.time() - start_time

            print(f"βœ“ BaseDataset initialization: {init_time:.3f}s for {len(dataset)} items")

            start_time = time.time()
            labels = dataset.get_labels()
            label_time = time.time() - start_time

            print(f"βœ“ Label extraction: {label_time:.3f}s for {len(labels)} labels")

            start_time = time.time()
            for i in range(min(10, len(dataset))):
                metadata = dataset.get_metadata(i)
            metadata_time = time.time() - start_time

            print(f"βœ“ Metadata extraction: {metadata_time:.3f}s for 10 items")

            assert init_time < 1.0
            assert label_time < 0.1

    # Run all tests
    tests = [
        ("BaseDataset Functionality", test_base_dataset),
        ("URLDataset with Mocked Requests", test_url_dataset_mock),
        ("URLDataset Fallback Behavior", test_url_dataset_fallback),
        ("CachedDataset Functionality", test_cached_dataset),
        ("Subset Creation", test_subset_creation),
        ("Bulk Download with Mocked Requests", test_download_images),
        ("Edge Cases and Error Handling", test_edge_cases),
        ("Performance Testing", test_performance),
    ]

    for test_name, test_func in tests:
        run_test(test_name, test_func)

    # Print summary
    print(f"\n{'=' * 60}")
    print(f"TEST SUMMARY")
    print(f"{'=' * 60}")
    print(f"Total tests: {len(tests)}")
    print(f"βœ… Passed: {tests_passed}")
    print(f"❌ Failed: {tests_failed}")

    if failed_tests:
        print(f"\nFailed tests:")
        for failure in failed_tests:
            print(f"  - {failure}")

    success = tests_failed == 0
    if success:
        print(f"\nπŸŽ‰ ALL TESTS PASSED! url_dataset.py is working correctly.")
    else:
        print(f"\n❌ SOME TESTS FAILED! Check the issues above.")

    return success


# Update the main section to include the comprehensive test option
if __name__ == "__main__":
    import sys

    if len(sys.argv) > 1 and sys.argv[1] == "--comprehensive":
        # Run comprehensive test
        run_comprehensive_test()
    else:
        # Run the quick test (your existing main code)
        # ... (your existing main code here)
        pass