diff --git a/docs/reference/triplet_store.md b/docs/reference/triplet_store.md index ad7a0645..ee5c24e6 100644 --- a/docs/reference/triplet_store.md +++ b/docs/reference/triplet_store.md @@ -182,7 +182,7 @@ for row in result.bindings: store = TripletStore( backend="rdf4j", endpoint="http://localhost:8080/rdf4j-server", - repository_id="semantica", # passed through **config + repository_id="semantica", # selects the remote repository ) ``` diff --git a/semantica/triplet_store/rdf4j_store.py b/semantica/triplet_store/rdf4j_store.py index c03b64a8..d788ab7b 100644 --- a/semantica/triplet_store/rdf4j_store.py +++ b/semantica/triplet_store/rdf4j_store.py @@ -28,7 +28,7 @@ License: MIT import re from typing import Any, Dict, List, Optional -from urllib.parse import urlparse +from urllib.parse import quote, urlparse import requests from rdflib import Graph, Literal @@ -68,6 +68,7 @@ class RDF4JStore: self.endpoint = endpoint.rstrip("/") self.repository_id = repository_id or config.get("repository_id", "default") + self._encoded_repository_id = quote(self.repository_id, safe="") self.username = config.get("username") self.password = config.get("password") self.timeout = config.get("timeout", 30) @@ -79,7 +80,7 @@ class RDF4JStore: """Connect to RDF4J server.""" try: # Test connection - test_url = f"{self.endpoint}/repositories/{self.repository_id}" + test_url = f"{self.endpoint}/repositories/{self._encoded_repository_id}" response = requests.get( test_url, timeout=self.timeout, @@ -100,11 +101,11 @@ class RDF4JStore: def _get_sparql_endpoint(self) -> str: """Get SPARQL query endpoint.""" - return f"{self.endpoint}/repositories/{self.repository_id}" + return f"{self.endpoint}/repositories/{self._encoded_repository_id}" def _get_update_endpoint(self) -> str: """Get SPARQL Update endpoint.""" - return f"{self.endpoint}/repositories/{self.repository_id}/statements" + return f"{self.endpoint}/repositories/{self._encoded_repository_id}/statements" def _is_construct_query(self, query: str) -> bool: """ @@ -163,7 +164,7 @@ class RDF4JStore: """ # RDF4J transaction support transaction_url = ( - f"{self.endpoint}/repositories/{self.repository_id}/transactions" + f"{self.endpoint}/repositories/{self._encoded_repository_id}/transactions" ) try: diff --git a/tests/triplet_store/test_rdf4j_store.py b/tests/triplet_store/test_rdf4j_store.py index 3f630b93..03a9b90b 100644 --- a/tests/triplet_store/test_rdf4j_store.py +++ b/tests/triplet_store/test_rdf4j_store.py @@ -23,6 +23,7 @@ CONSTRUCT_QUERY = "CONSTRUCT { ?s ?p ?o } WHERE { ?s ?p ?o }" class TestRDF4JStoreInitialization(unittest.TestCase): + def test_explicit_repository_id_selects_repository(self): response = MagicMock(status_code=200) @@ -42,6 +43,49 @@ class TestRDF4JStoreInitialization(unittest.TestCase): auth=None, ) + def test_repository_id_is_encoded_as_a_single_url_path_segment(self): + response = MagicMock(status_code=200) + + with patch( + "semantica.triplet_store.rdf4j_store.requests.get", + return_value=response, + ) as mock_get: + store = RDF4JStore( + endpoint="http://localhost:8080/rdf4j-server", + repository_id="team/repo ?#", + ) + + self.assertEqual(store.repository_id, "team/repo ?#") + mock_get.assert_called_once_with( + "http://localhost:8080/rdf4j-server/repositories/team%2Frepo%20%3F%23", + timeout=30, + auth=None, + ) + self.assertEqual( + store._get_sparql_endpoint(), + "http://localhost:8080/rdf4j-server/repositories/team%2Frepo%20%3F%23", + ) + self.assertEqual( + store._get_update_endpoint(), + "http://localhost:8080/rdf4j-server/repositories/" + "team%2Frepo%20%3F%23/statements", + ) + + transaction_response = MagicMock() + transaction_response.headers = {"Location": "/transactions/tx-1"} + with patch( + "semantica.triplet_store.rdf4j_store.requests.post", + return_value=transaction_response, + ) as mock_post: + self.assertEqual(store.begin_transaction(), "tx-1") + + mock_post.assert_called_once_with( + "http://localhost:8080/rdf4j-server/repositories/" + "team%2Frepo%20%3F%23/transactions", + timeout=30, + auth=None, + ) + class TestRDF4JStoreIsConstructQuery(unittest.TestCase): def test_detects_uppercase(self):