File size: 6,335 Bytes
343140f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Unit tests for the ProTox 3.0 client (app.tools.protox).

These mock the upstream HTTP calls so no network access is required.
Covers: request validation, enqueue/retrieve/fetch flow, tolerant TSV
parsing, and graceful error mapping for quota / rate-limit / outage.
"""

import pytest
import httpx

from app.tools.protox import (
    ALL_MODELS,
    DEFAULT_MODELS,
    ProToxError,
    _parse_tsv,
    predict_toxicity,
)


class FakeResponse:
    def __init__(self, status_code: int, text: str = ""):
        self.status_code = status_code
        self.text = text

    def raise_for_status(self):
        if self.status_code >= 400:
            raise httpx.HTTPStatusError(
                f"{self.status_code}", request=httpx.Request("GET", "http://x"), response=self
            )


TOX_CLASS_CSV = (
    'LD50\ttox_class\tsimilarity\n'
    '100\t3\t78.9\n'
)
MODELS_CSV = (
    'Target\tPrediction\tProbability\n'
    'dili\t1\t0.93\n'
    'carcino\t0\t0.12\n'
)
TARGETS_CSV = (
    'Target\tPrediction\tProbability\n'
    'Membrane alpha-2 receptor\t1\t0.81\n'
)


def make_client_side_effect(enqueue_status=200, enqueue_text="abc123"):
    """Return an httpx.AsyncClient whose post/get follow the ProTox flow."""

    async def _post(url, data=None, **kwargs):
        if url.endswith("api_enqueue.php"):
            return FakeResponse(enqueue_status, enqueue_text)
        if url.endswith("api_retrieve.php"):
            # Non-empty body on first poll => ready immediately.
            return FakeResponse(200, "done")
        raise AssertionError(f"unexpected POST url: {url}")

    async def _get(url, **kwargs):
        if url.endswith("_tox_class.csv"):
            return FakeResponse(200, TOX_CLASS_CSV)
        if url.endswith("_result.csv"):
            return FakeResponse(200, MODELS_CSV)
        if url.endswith("_tox_targets.csv"):
            return FakeResponse(200, TARGETS_CSV)
        raise AssertionError(f"unexpected GET url: {url}")

    return _post, _get


class FakeAsyncClient:
    def __init__(self, post_fn, get_fn):
        self._post = post_fn
        self._get = get_fn

    async def __aenter__(self):
        return self

    async def __aexit__(self, *exc):
        return False

    async def post(self, url, data=None, **kwargs):
        return await self._post(url, data=data, **kwargs)

    async def get(self, url, **kwargs):
        return await self._get(url, **kwargs)


def test_parse_tsv_normalizes_headers():
    rows = _parse_tsv('Target\tPrediction\tProbability\nDILI\t1\t0.9\n')
    assert rows == [{"target": "DILI", "prediction": "1", "probability": "0.9"}]


def test_parse_tsv_handles_quoted_headers_and_blank_lines():
    rows = _parse_tsv('"Target"\tPrediction\nx\t1\n\n')
    assert rows == [{"target": "x", "prediction": "1"}]


def test_parse_tsv_empty():
    assert _parse_tsv("") == []
    assert _parse_tsv("no tabs here\n") == []


def test_request_validation():
    with pytest.raises(ProToxError):
        asyncio_run(predict_toxicity())
    with pytest.raises(ProToxError):
        asyncio_run(predict_toxicity(smiles="CCO", name="ethanol"))


def asyncio_run(coro):
    import asyncio
    return asyncio.run(coro)


def test_successful_flow(monkeypatch):
    post, get = make_client_side_effect()
    monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: FakeAsyncClient(post, get))

    result = asyncio_run(predict_toxicity(smiles="CCO"))

    assert result["input"] == "CCO"
    assert result["input_type"] == "smiles"
    assert result["task_id"] == "abc123"
    assert result["acute_toxicity"]["ld50"] == "100"
    assert result["acute_toxicity"]["tox_class"] == "3"
    assert result["model_results"] == [
        {"target": "dili", "prediction": "1", "probability": "0.93"},
        {"target": "carcino", "prediction": "0", "probability": "0.12"},
    ]
    assert result["toxicity_targets"][0]["target"] == "Membrane alpha-2 receptor"
    assert result["methodology"]["tier"] == "3a"


def test_name_input_flow(monkeypatch):
    post, get = make_client_side_effect()
    monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: FakeAsyncClient(post, get))

    result = asyncio_run(predict_toxicity(name="aspirin", models="acute_tox dili"))
    assert result["input_type"] == "name"
    assert result["requested_models"] == "acute_tox dili"


def test_quota_exceeded_raises(monkeypatch):
    post, get = make_client_side_effect(enqueue_status=403)
    monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: FakeAsyncClient(post, get))

    with pytest.raises(ProToxError) as exc:
        asyncio_run(predict_toxicity(smiles="CCO"))
    assert "quota" in str(exc.value).lower()


def test_rate_limited_raises(monkeypatch):
    post, get = make_client_side_effect(enqueue_status=429)
    monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: FakeAsyncClient(post, get))

    with pytest.raises(ProToxError) as exc:
        asyncio_run(predict_toxicity(smiles="CCO"))
    assert "throttl" in str(exc.value).lower()


def test_outage_raises_graceful(monkeypatch):
    post, get = make_client_side_effect(enqueue_status=404)
    monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: FakeAsyncClient(post, get))

    with pytest.raises(ProToxError) as exc:
        asyncio_run(predict_toxicity(smiles="CCO"))
    assert "unavailable" in str(exc.value) or "404" in str(exc.value)


def test_empty_task_id_raises(monkeypatch):
    post, get = make_client_side_effect(enqueue_status=200, enqueue_text=" ")
    monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: FakeAsyncClient(post, get))

    with pytest.raises(ProToxError) as exc:
        asyncio_run(predict_toxicity(smiles="CCO"))
    assert "task id" in str(exc.value).lower()


def test_model_constant_sanity():
    assert ALL_MODELS.split()[0] == "dili"
    assert "CYP3A4" in ALL_MODELS
    assert "acute_tox" in DEFAULT_MODELS
    assert "tox_targets" in DEFAULT_MODELS


@pytest.mark.parametrize("models", [None, "acute_tox", "ALL_MODELS"])
def test_model_arg_variants(monkeypatch, models):
    post, get = make_client_side_effect()
    monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: FakeAsyncClient(post, get))

    kwargs = {"smiles": "CCO"}
    if models:
        kwargs["models"] = models
    result = asyncio_run(predict_toxicity(**kwargs))
    assert result["requested_models"]