|
1 | 1 | import pytest |
| 2 | +from unittest import mock |
2 | 3 | import numpy as np |
3 | 4 | import pandas as pd |
4 | 5 |
|
|
14 | 15 | ) |
15 | 16 | from mlflow.pyfunc import PyFuncModel |
16 | 17 | from mlflow.models.signature import ModelSignature |
| 18 | +from mlflow.pyfunc.scoring_server import CONTENT_TYPE_CSV, CONTENT_TYPE_JSON |
17 | 19 |
|
18 | 20 | from mlserver_mlflow import MLflowRuntime |
19 | 21 | from mlserver_mlflow.codecs import TensorDictCodec |
@@ -188,3 +190,65 @@ async def test_metadata(runtime: MLflowRuntime, model_signature: ModelSignature) |
188 | 190 |
|
189 | 191 | assert metadata.parameters is not None |
190 | 192 | assert metadata.parameters.content_type == PandasCodec.ContentType |
| 193 | + |
| 194 | + |
| 195 | +@pytest.mark.parametrize( |
| 196 | + "input, expected", |
| 197 | + [ |
| 198 | + # works with params: |
| 199 | + ( |
| 200 | + ['{"instances": [1, 2, 3], "params": {"foo": "bar"}}', CONTENT_TYPE_JSON], |
| 201 | + {"data": {"foo": [1, 2, 3]}, "params": {"foo": "bar"}}, |
| 202 | + ), |
| 203 | + ( |
| 204 | + [ |
| 205 | + '{"inputs": [1, 2, 3], "params": {"foo": "bar"}}', |
| 206 | + CONTENT_TYPE_JSON, |
| 207 | + ], |
| 208 | + {"data": {"foo": [1, 2, 3]}, "params": {"foo": "bar"}}, |
| 209 | + ), |
| 210 | + ( |
| 211 | + [ |
| 212 | + '{"inputs": {"foo": [1, 2, 3]}, "params": {"foo": "bar"}}', |
| 213 | + CONTENT_TYPE_JSON, |
| 214 | + ], |
| 215 | + {"data": {"foo": [1, 2, 3]}, "params": {"foo": "bar"}}, |
| 216 | + ), |
| 217 | + ( |
| 218 | + [ |
| 219 | + '{"dataframe_split": {"columns": ["foo"], "data": [1, 2, 3]}, "params": {"foo": "bar"}}', |
| 220 | + CONTENT_TYPE_JSON, |
| 221 | + ], |
| 222 | + {"data": {"foo": [1, 2, 3]}, "params": {"foo": "bar"}}, |
| 223 | + ), |
| 224 | + ( |
| 225 | + [ |
| 226 | + '{"dataframe_records": [{"foo": 1}, {"foo": 2}, {"foo": 3}], "params": {"foo": "bar"}}', |
| 227 | + CONTENT_TYPE_JSON, |
| 228 | + ], |
| 229 | + {"data": {"foo": [1, 2, 3]}, "params": {"foo": "bar"}}, |
| 230 | + ), |
| 231 | + ( |
| 232 | + ["foo\n1\n2\n3\n", CONTENT_TYPE_CSV], |
| 233 | + {"data": {"foo": [1, 2, 3]}, "params": None}, |
| 234 | + ), |
| 235 | + # works without params: |
| 236 | + ( |
| 237 | + ['{"instances": [1, 2, 3]}', CONTENT_TYPE_JSON], |
| 238 | + {"data": {"foo": [1, 2, 3]}, "params": None}, |
| 239 | + ), |
| 240 | + ], |
| 241 | +) |
| 242 | +async def test_invocation_with_params( |
| 243 | + runtime: MLflowRuntime, |
| 244 | + input: list, |
| 245 | + expected: dict, |
| 246 | +): |
| 247 | + with mock.patch.object( |
| 248 | + runtime._model, "predict", return_value=[1, 2, 3] |
| 249 | + ) as predict_mock: |
| 250 | + await runtime.invocations(*input) |
| 251 | + np.testing.assert_array_equal( |
| 252 | + predict_mock.call_args[0][0].get("foo"), expected["data"]["foo"] |
| 253 | + ) |
| 254 | + assert predict_mock.call_args.kwargs["params"] == expected["params"] |
0 commit comments