test_model_provider_registry.py 2.4 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970
  1. """Tests for the model-provider registry.
  2. Pins the routing seam that a future *shared* import API will use: pasted URLs
  3. go through ``find_for_url`` and land on the provider that owns them.
  4. """
  5. from __future__ import annotations
  6. import pytest
  7. from backend.app.services.model_providers import makerworld_provider, registry
  8. from backend.app.services.model_providers.base import ModelProvider
  9. from backend.app.services.model_providers.registry import ModelProviderRegistry
  10. class _DummyProvider(ModelProvider):
  11. source_type = "dummy"
  12. display_name = "Dummy"
  13. async def build_service(self, *, db, user, api_key_owner=None, client=None):
  14. raise NotImplementedError
  15. def parse_url(self, url):
  16. raise NotImplementedError
  17. def canonical_url(self, ref):
  18. raise NotImplementedError
  19. class TestAppRegistry:
  20. """The app-wide singleton auto-registers MakerWorld on import."""
  21. def test_makerworld_is_registered(self):
  22. assert registry.get("makerworld") is makerworld_provider
  23. assert registry.get("makerworld").display_name == "MakerWorld"
  24. def test_unknown_source_type_raises_keyerror(self):
  25. with pytest.raises(KeyError):
  26. registry.get("thingiverse")
  27. def test_find_for_url_routes_makerworld_urls(self):
  28. provider = registry.find_for_url("https://makerworld.com/en/models/1400373#profileId-1452154")
  29. assert provider is makerworld_provider
  30. def test_find_for_url_returns_none_for_foreign_hosts(self):
  31. assert registry.find_for_url("https://thingiverse.com/thing/123") is None
  32. assert registry.find_for_url("") is None
  33. assert registry.find_for_url(None) is None # type: ignore[arg-type]
  34. class TestModelProviderRegistry:
  35. def test_register_is_idempotent_per_instance(self):
  36. reg = ModelProviderRegistry()
  37. provider = _DummyProvider()
  38. reg.register(provider)
  39. reg.register(provider)
  40. assert reg.all() == (provider,)
  41. def test_register_duplicate_source_type_rejected(self):
  42. reg = ModelProviderRegistry()
  43. reg.register(_DummyProvider())
  44. with pytest.raises(ValueError):
  45. reg.register(_DummyProvider())
  46. def test_all_returns_registered_providers(self):
  47. reg = ModelProviderRegistry()
  48. provider = _DummyProvider()
  49. reg.register(provider)
  50. assert provider in reg.all()
  51. assert len(reg.all()) == 1