Text Classification
Keras
English
Azerbaijani
prompt-injection
security
llm-security
document-security
retvec
cnn
tensorflow
fastapi
Eval Results (legacy)
Instructions to use MegrurNiftiyev/MyGuard-Prompt-Injection-Detector with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Keras
How to use MegrurNiftiyev/MyGuard-Prompt-Injection-Detector with Keras:
# Available backend options are: "jax", "torch", "tensorflow". import os os.environ["KERAS_BACKEND"] = "jax" import keras model = keras.saving.load_model("hf://MegrurNiftiyev/MyGuard-Prompt-Injection-Detector") - Notebooks
- Google Colab
- Kaggle
File size: 6,413 Bytes
215f97f | 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 | """
Tests for the /classify endpoint with fullText schema & chunk prediction.
"""
import pytest
from unittest.mock import AsyncMock, patch
from httpx import AsyncClient, ASGITransport
from app.main import app
from app.ml.serving.registry import DummyModel
@pytest.fixture
def auth_headers():
"""Valid internal service auth headers."""
return {"X-Internal-Token": "test-secret"}
@pytest.mark.asyncio
async def test_classify_returns_prediction(auth_headers):
"""POST /classify should return a valid ClassifyResponse for fullText."""
dummy = DummyModel()
with patch(
"app.api.routes.classify.load_active_model",
new_callable=AsyncMock,
return_value=dummy,
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/analyze-injection",
json={
"documentId": "doc-123",
"fullText": "This is a normal corporate document with standard operational content.",
},
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
assert data["label"] in ("safe", "suspicious", "injection")
assert 0.0 <= data["confidence"] <= 1.0
@pytest.mark.asyncio
async def test_classify_without_document_id(auth_headers):
"""POST /analyze-injection should work even if documentId is omitted."""
dummy = DummyModel()
with patch(
"app.api.routes.classify.load_active_model",
new_callable=AsyncMock,
return_value=dummy,
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/analyze-injection",
json={
"fullText": "This is a normal corporate document with standard operational content.",
},
headers=auth_headers,
)
assert response.status_code == 200
data = response.json()
assert data["label"] in ("safe", "suspicious", "injection")
assert 0.0 <= data["confidence"] <= 1.0
@pytest.mark.skip(reason="Token check temporarily disabled for local dev testing")
@pytest.mark.asyncio
async def test_classify_rejects_missing_auth():
"""POST /classify without X-Internal-Token should return 401 or 422."""
dummy = DummyModel()
with patch(
"app.api.routes.classify.load_active_model",
new_callable=AsyncMock,
return_value=dummy,
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/analyze-injection",
json={
"documentId": "doc-123",
"fullText": "Sample text for testing authentication validation.",
},
)
assert response.status_code in (401, 422)
@pytest.mark.skip(reason="Token check temporarily disabled for local dev testing")
@pytest.mark.asyncio
async def test_classify_rejects_wrong_token():
"""POST /classify with wrong token should return 401."""
dummy = DummyModel()
with patch(
"app.api.routes.classify.load_active_model",
new_callable=AsyncMock,
return_value=dummy,
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/analyze-injection",
json={
"documentId": "doc-123",
"fullText": "Sample text for testing invalid token handling.",
},
headers={"X-Internal-Token": "wrong-secret"},
)
assert response.status_code == 401
@pytest.mark.asyncio
async def test_classify_rejects_insufficient_text(auth_headers):
"""Verify that fullText under 5 words raises 503 insufficient_text."""
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/analyze-injection",
json={
"documentId": "doc-short",
"fullText": "One two three four", # 4 words
},
headers=auth_headers,
)
assert response.status_code == 503
assert response.json()["message"] == "insufficient_text"
@pytest.mark.asyncio
async def test_classify_rejects_extra_legacy_fields(auth_headers):
"""Verify that extra legacy fields (text, ocrText, hiddenText) are rejected (422)."""
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
response = await client.post(
"/analyze-injection",
json={
"documentId": "doc-legacy",
"fullText": "This is valid text containing enough words for test.",
"text": "Legacy text field that should be forbidden",
},
headers=auth_headers,
)
assert response.status_code == 422
@pytest.mark.asyncio
async def test_classify_passes_raw_full_text(auth_headers):
"""Verify raw fullText is passed directly to run_prediction."""
dummy = DummyModel()
captured_texts = []
def capturing_run_prediction(model, text):
captured_texts.append(text)
return ("safe", 0.99)
raw_input_text = " Hello WORLD\nLine two of document.\nLine three of document text."
with patch(
"app.api.routes.classify.load_active_model",
new_callable=AsyncMock,
return_value=dummy,
), patch(
"app.api.routes.classify.run_prediction",
side_effect=capturing_run_prediction,
):
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as client:
res = await client.post(
"/analyze-injection",
json={
"documentId": "doc-raw",
"fullText": raw_input_text,
},
headers=auth_headers,
)
assert res.status_code == 200
assert captured_texts[0] == raw_input_text
|