File size: 3,845 Bytes
5752a28
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import os
import sys
import unittest


BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if BACKEND_DIR not in sys.path:
    sys.path.insert(0, BACKEND_DIR)

from server import app, model_service


class ServerTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        app.config.update(TESTING=True)
        cls.client = app.test_client()

    def test_grouped_metrics_have_no_page_leakage(self):
        response = self.client.get("/api/metrics")
        self.assertEqual(response.status_code, 200)
        metrics = response.get_json()
        self.assertEqual(metrics["groupOverlap"], 0)
        self.assertIn("page_id", metrics["splitMethod"])
        self.assertTrue(metrics["calibrated"])
        self.assertGreater(metrics["charFeatureCount"], 0)

    def test_single_text_classification_is_calibrated(self):
        response = self.client.post(
            "/api/analyze-text",
            json={"text": "Only two investment slots left. Act now."},
        )
        self.assertEqual(response.status_code, 200)
        result = response.get_json()
        self.assertTrue(result["isDarkPattern"])
        self.assertIn(result["prediction"], model_service.classes)
        self.assertTrue(result["calibrated"])
        self.assertGreaterEqual(result["confidence"], 0)
        self.assertLessEqual(result["confidence"], 100)

    def test_batch_classification_preserves_ids(self):
        response = self.client.post(
            "/api/analyze-texts",
            json={
                "items": [
                    {"id": "cta", "text": "Hurry, only 2 left!"},
                    {
                        "id": "statement",
                        "text": "Your monthly statement is ready.",
                    },
                ]
            },
        )
        self.assertEqual(response.status_code, 200)
        payload = response.get_json()
        self.assertEqual(payload["analyzed"], 2)
        self.assertEqual(
            [result["id"] for result in payload["results"]],
            ["cta", "statement"],
        )

    def test_remote_image_fetch_is_blocked(self):
        response = self.client.post(
            "/api/analyze",
            json={"imageUrl": "https://example.com/private-image.png"},
        )
        self.assertEqual(response.status_code, 400)
        self.assertIn(
            "Remote image URLs are disabled", response.get_json()["error"]
        )

    def test_cors_allows_extension_and_rejects_unrelated_sites(self):
        allowed = self.client.options(
            "/api/analyze-texts",
            headers={
                "Origin": "chrome-extension://abcdefghijklmnopabcdefghijklmnop",
                "Access-Control-Request-Method": "POST",
            },
        )
        self.assertEqual(
            allowed.headers.get("Access-Control-Allow-Origin"),
            "chrome-extension://abcdefghijklmnopabcdefghijklmnop",
        )

        denied = self.client.options(
            "/api/analyze-texts",
            headers={
                "Origin": "https://unrelated.example",
                "Access-Control-Request-Method": "POST",
            },
        )
        self.assertIsNone(denied.headers.get("Access-Control-Allow-Origin"))

    def test_sample_ocr_response_schema(self):
        response = self.client.post(
            "/api/analyze",
            json={
                "imageUrl": (
                    "http://127.0.0.1:8000/api/samples/checkout_urgency.png"
                )
            },
        )
        self.assertEqual(response.status_code, 200)
        payload = response.get_json()
        self.assertGreaterEqual(payload["ocrMinimumConfidence"], 0)
        self.assertIsInstance(payload["darkPatterns"], list)
        self.assertIn("complianceReport", payload)


if __name__ == "__main__":
    unittest.main()