diff --git a/pycsw/core/repository.py b/pycsw/core/repository.py index 007b7ead8..9f890c13b 100644 --- a/pycsw/core/repository.py +++ b/pycsw/core/repository.py @@ -435,6 +435,14 @@ def query(self, constraint, sortby=None, typenames=None, LOGGER.debug('No constraint detected') query = self.session.query(self.dataset) + # filter by typename when specific typenames requested + _csw_generic = {'csw:Record', 'csw30:Record'} + if typenames and not any(t in _csw_generic for t in typenames): + typename_col = self.context.md_core_model['mappings']['pycsw:Typename'] + query = query.filter( + getattr(self.dataset, typename_col).in_(typenames) + ) + total = self._get_repo_filter(query).count() if util.ranking_pass: # apply spatial ranking diff --git a/pycsw/ogc/csw/csw2.py b/pycsw/ogc/csw/csw2.py index 56b9547f7..a0c9b24ee 100644 --- a/pycsw/ogc/csw/csw2.py +++ b/pycsw/ogc/csw/csw2.py @@ -733,7 +733,7 @@ def getrecords(self): self.parent.kvp['constraint']['type'] = 'filter' cql = cql2fes(tmp, self.parent.context.namespaces, fes_version='1.0') self.parent.kvp['constraint']['where'], self.parent.kvp['constraint']['values'] = fes1.parse(cql, - self.parent.repository.queryables['_all'], self.parent.repository.dbtype, + self._scoped_queryables(self.parent.kvp.get('typenames')), self.parent.repository.dbtype, self.parent.context.namespaces, self.parent.orm, self.parent.language['text'], self.parent.repository.fts) self.parent.kvp['constraint']['_dict'] = xml2dict(etree.tostring(cql), self.parent.context.namespaces) except Exception as err: @@ -754,7 +754,7 @@ def getrecords(self): self.parent.kvp['constraint']['type'] = 'filter' self.parent.kvp['constraint']['where'], self.parent.kvp['constraint']['values'] = \ fes1.parse(doc, - self.parent.repository.queryables['_all'], + self._scoped_queryables(self.parent.kvp.get('typenames')), self.parent.repository.dbtype, self.parent.context.namespaces, self.parent.orm, self.parent.language['text'], self.parent.repository.fts) self.parent.kvp['constraint']['_dict'] = xml2dict(etree.tostring(doc), self.parent.context.namespaces) @@ -1567,7 +1567,21 @@ def _write_record(self, recobj, queryables): record.append(bboxel) return record - def _parse_constraint(self, element): + def _scoped_queryables(self, typenames=None): + ''' Scope queryables to the requested typenames so profiles sharing a + property name resolve to the correct column ''' + csw_generic = {'csw:Record', 'csw30:Record'} + tnames = typenames if isinstance(typenames, list) else (typenames.split() if typenames else []) + if not tnames or any(t in csw_generic for t in tnames): + return self.parent.repository.queryables['_all'] + scoped = dict(self.parent.repository.queryables['_all']) + for tname in tnames: + if tname in self.parent.context.model['typenames']: + for qgroup in self.parent.context.model['typenames'][tname]['queryables']: + scoped.update(self.parent.repository.queryables.get(qgroup, {})) + return scoped + + def _parse_constraint(self, element, typenames=None): ''' Parse csw:Constraint ''' query = {} @@ -1578,7 +1592,7 @@ def _parse_constraint(self, element): try: query['type'] = 'filter' query['where'], query['values'] = fes1.parse(tmp, - self.parent.repository.queryables['_all'], self.parent.repository.dbtype, + self._scoped_queryables(typenames), self.parent.repository.dbtype, self.parent.context.namespaces, self.parent.orm, self.parent.language['text'], self.parent.repository.fts) query['_dict'] = xml2dict(etree.tostring(tmp), self.parent.context.namespaces) except Exception as err: @@ -1592,7 +1606,7 @@ def _parse_constraint(self, element): query['type'] = 'filter' cql = cql2fes(tmp.text, self.parent.context.namespaces, fes_version='1.0') query['where'], query['values'] = fes1.parse(cql, - self.parent.repository.queryables['_all'], self.parent.repository.dbtype, + self._scoped_queryables(typenames), self.parent.repository.dbtype, self.parent.context.namespaces, self.parent.orm, self.parent.language['text'], self.parent.repository.fts) query['_dict'] = xml2dict(etree.tostring(cql), self.parent.context.namespaces) except Exception as err: @@ -1780,7 +1794,8 @@ def parse_postdata(self, postdata): self.parent.context.namespaces)) if tmp is not None: - request['constraint'] = self._parse_constraint(tmp) + request['constraint'] = self._parse_constraint( + tmp, typenames=request.get('typenames')) if isinstance(request['constraint'], str): # parse error return 'Invalid Constraint: %s' % request['constraint'] else: diff --git a/pycsw/ogc/csw/csw3.py b/pycsw/ogc/csw/csw3.py index d095d7de0..bac015703 100644 --- a/pycsw/ogc/csw/csw3.py +++ b/pycsw/ogc/csw/csw3.py @@ -761,7 +761,7 @@ def getrecords(self): self.parent.kvp['constraint']['type'] = 'filter' cql = cql2fes(tmp, self.parent.context.namespaces, fes_version='1.0') self.parent.kvp['constraint']['where'], self.parent.kvp['constraint']['values'] = fes1.parse(cql, - self.parent.repository.queryables['_all'], self.parent.repository.dbtype, + self._scoped_queryables(self.parent.kvp.get('typenames')), self.parent.repository.dbtype, self.parent.context.namespaces, self.parent.orm, self.parent.language['text'], self.parent.repository.fts) self.parent.kvp['constraint']['_dict'] = xml2dict(etree.tostring(cql), self.parent.context.namespaces) except Exception as err: @@ -782,7 +782,7 @@ def getrecords(self): self.parent.kvp['constraint']['type'] = 'filter' self.parent.kvp['constraint']['where'], self.parent.kvp['constraint']['values'] = \ fes2.parse(doc, - self.parent.repository.queryables['_all'], + self._scoped_queryables(self.parent.kvp.get('typenames')), self.parent.repository.dbtype, self.parent.context.namespaces, self.parent.orm, self.parent.language['text'], self.parent.repository.fts) self.parent.kvp['constraint']['_dict'] = xml2dict(etree.tostring(doc), self.parent.context.namespaces) @@ -1641,7 +1641,21 @@ def _write_record(self, recobj, queryables): return record - def _parse_constraint(self, element): + def _scoped_queryables(self, typenames=None): + ''' Scope queryables to the requested typenames so profiles sharing a + property name resolve to the correct column ''' + csw_generic = {'csw:Record', 'csw30:Record'} + tnames = typenames if isinstance(typenames, list) else (typenames.split() if typenames else []) + if not tnames or any(t in csw_generic for t in tnames): + return self.parent.repository.queryables['_all'] + scoped = dict(self.parent.repository.queryables['_all']) + for tname in tnames: + if tname in self.parent.context.model['typenames']: + for qgroup in self.parent.context.model['typenames'][tname]['queryables']: + scoped.update(self.parent.repository.queryables.get(qgroup, {})) + return scoped + + def _parse_constraint(self, element, typenames=None): ''' Parse csw:Constraint ''' query = {} @@ -1652,7 +1666,7 @@ def _parse_constraint(self, element): try: query['type'] = 'filter' query['where'], query['values'] = fes2.parse(tmp, - self.parent.repository.queryables['_all'], self.parent.repository.dbtype, + self._scoped_queryables(typenames), self.parent.repository.dbtype, self.parent.context.namespaces, self.parent.orm, self.parent.language['text'], self.parent.repository.fts) query['_dict'] = xml2dict(etree.tostring(tmp), self.parent.context.namespaces) except Exception as err: @@ -1666,7 +1680,7 @@ def _parse_constraint(self, element): query['type'] = 'filter' cql = cql2fes(tmp.text, self.parent.context.namespaces, fes_version='2.0') query['where'], query['values'] = fes2.parse(cql, - self.parent.repository.queryables['_all'], self.parent.repository.dbtype, + self._scoped_queryables(typenames), self.parent.repository.dbtype, self.parent.context.namespaces, self.parent.orm, self.parent.language['text'], self.parent.repository.fts) query['_dict'] = xml2dict(etree.tostring(cql), self.parent.context.namespaces) except Exception as err: @@ -1846,7 +1860,8 @@ def parse_postdata(self, postdata): self.parent.context.namespaces)) if tmp is not None: - request['constraint'] = self._parse_constraint(tmp) + request['constraint'] = self._parse_constraint( + tmp, typenames=request.get('typenames')) if isinstance(request['constraint'], str): # parse error return 'Invalid Constraint: %s' % request['constraint'] else: diff --git a/tests/unittests/test_repository.py b/tests/unittests/test_repository.py index 6b00221f1..8300793bb 100644 --- a/tests/unittests/test_repository.py +++ b/tests/unittests/test_repository.py @@ -31,6 +31,7 @@ import pytest from pycsw.core import repository +from pycsw.core.config import StaticContext pytestmark = pytest.mark.unit @@ -64,3 +65,65 @@ def test_query_spatial(data, input_, predicate, distance, expected): distance=distance ) assert result == expected + + +@pytest.fixture +def multiprofile_repo(tmp_path): + """Repository seeded with records of two different typenames, as happens + when multiple CSW profiles share a single table.""" + database = 'sqlite:///%s' % (tmp_path / 'records.db') + repository.setup(database, 'records') + + context = StaticContext() + repo = repository.Repository(database, context, table='records') + + fixtures = [ + ('rec-generic', 'csw:Record'), + ('rec-profile-a', 'foo:RecordA'), + ('rec-profile-a-2', 'foo:RecordA'), + ('rec-profile-b', 'foo:RecordB'), + ] + for identifier, typename in fixtures: + record = repo.dataset( + identifier=identifier, + typename=typename, + schema='http://www.opengis.net/cat/csw/2.0.2', + mdsource='local', + insert_date='2024-01-01', + xml='', + anytext=identifier, + metadata_type='application/xml', + ) + repo.insert(record, 'local', '2024-01-01') + + return repo + + +def test_query_filters_by_typename(multiprofile_repo): + """A GetRecords with a specific (non-generic) typeNames must only return + rows whose typename column matches.""" + total, records = multiprofile_repo.query( + constraint={}, typenames=['foo:RecordA']) + + assert total == '2' + assert {r.identifier for r in records} == {'rec-profile-a', + 'rec-profile-a-2'} + assert all(r.typename == 'foo:RecordA' for r in records) + + +def test_query_generic_typename_returns_all(multiprofile_repo): + """The CSW wildcard type csw:Record must not filter by typename, i.e. + single-profile / default behaviour is preserved.""" + total, records = multiprofile_repo.query( + constraint={}, typenames=['csw:Record']) + + assert total == '4' + assert len(records) == 4 + + +def test_query_no_typename_returns_all(multiprofile_repo): + """Omitting typenames entirely preserves the pre-existing behaviour.""" + total, records = multiprofile_repo.query(constraint={}) + + assert total == '4' + assert len(records) == 4