Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 8 additions & 3 deletions ado/core/samplestore/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -369,7 +369,9 @@ def from_configuration(
f"Copying {sample_store_source.numberOfEntities} entities from "
f"{sample_store_source.identifier} to {sample_store.identifier}"
)
sample_store.add_external_entities(sample_store_source.entities)
sample_store.add_external_entities(
sample_store_source.get_entities(require_measurements=True)
)

return sample_store

Expand Down Expand Up @@ -517,13 +519,16 @@ def addMeasurement(
def entityWithIdentifier(
self, entityIdentifier: str
) -> Entity | None: # pragma: nocover
pass
"""Deprecated: use :meth:`get_entities` instead.

Returns entity if it is in the store, otherwise returns ``None``.
"""

@abc.abstractmethod
def entities_with_identifiers(
self, entity_identifiers: set[str] | list[str]
) -> list[Entity]:
"""Fetch the entities given by entity_identifiers.
"""Deprecated: use :meth:`get_entities` instead.

Args:
entity_identifiers: Set or list of entity identifiers to fetch
Expand Down
4 changes: 3 additions & 1 deletion ado/core/samplestore/csv.py
Original file line number Diff line number Diff line change
Expand Up @@ -307,7 +307,9 @@ def __init__(
for result in measurement_results:
entity.add_measurement_result(result)

self._entity_ids = [e.identifier for e in self.entities]
self._entity_ids = [
e.identifier for e in self.get_entities(require_measurements=False)
]

@property
def config(self) -> CSVSampleStoreDescription:
Expand Down
195 changes: 38 additions & 157 deletions ado/core/samplestore/sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,7 +125,9 @@ def from_csv(
storageLocation=storeConfiguration,
parameters={},
)
sql_sample_store.add_external_entities(csv_sample_store.entities)
sql_sample_store.add_external_entities(
csv_sample_store.get_entities(require_measurements=True)
)

return sql_sample_store

Expand Down Expand Up @@ -222,7 +224,7 @@ def experimentCatalog(

# Step 2: load those entities (with their measurement results) via the
# existing method which handles caching, decoding, and result attachment.
entities = self.entities_with_identifiers(entity_ids)
entities = self.get_entities(identifiers=entity_ids, require_measurements=True)

# Step 3: build the catalog from the fully-loaded entities.
experiments: dict[str, Experiment] = {}
Expand Down Expand Up @@ -518,7 +520,14 @@ def _fetch_entities(self, entity_ids: set[str] | None = None) -> dict[str, Entit
f"Fetched {len(entities)} entities"
+ (f" (filtered from {len(entity_ids)} requested)" if entity_ids else "")
)
# Always merge fetched entities into the cache.
# When merging, invalidate measurement-loaded status for any entity
# whose cached object is being replaced — the fresh DB representation
# does not carry measurement results, so they must be re-fetched.
overwritten_ids = set(entities).intersection(
self._entities_with_measurements_loaded
)
if overwritten_ids:
self._entities_with_measurements_loaded.difference_update(overwritten_ids)
self._entities.update(entities)
return entities

Expand Down Expand Up @@ -734,107 +743,22 @@ def refresh(self, force_fetch_all_entities: bool = False) -> tuple[int, int]:
def entities_with_identifiers(
self, entity_identifiers: set[str] | list[str]
) -> list[Entity]:
"""Efficiently fetch entities by their identifiers without loading all entities.

This method queries only the specified entities from the database, making it
much more efficient than loading all entities and filtering in Python.
"""Deprecated: use :meth:`get_entities` instead.

Args:
entity_identifiers: Set or list of entity identifiers to fetch

Returns:
List of Entity objects matching the provided identifiers
"""
if not entity_identifiers:
return []

# Convert to set for deduplication and efficient lookup
entity_ids_set = (
set(entity_identifiers)
if isinstance(entity_identifiers, list)
else entity_identifiers
warnings.warn(
"entities_with_identifiers is deprecated, use get_entities instead.",
DeprecationWarning,
stacklevel=2,
)

# Partition into cached and uncached IDs
cached_keys = (
entity_ids_set.intersection(self._entities.keys())
if self._entities
else set()
return self.get_entities(
identifiers=set(entity_identifiers), require_measurements=True
)
uncached_ids = entity_ids_set.difference(cached_keys)
cached_entities = [self._entities[k] for k in cached_keys]

# All requested entities were already cached
if not uncached_ids:
return cached_entities

# Query database only for the uncached entities
# Use SQLAlchemy's expanding bindparam for IN clause
# This automatically handles the parameter expansion for the IN clause
query = sqlalchemy.text(f"""
SELECT ent.identifier, ent.representation, res.data
FROM {self._tablename} ent
LEFT OUTER JOIN {self._tablename}_measurement_results res ON res.entity_id = ent.identifier
WHERE ent.identifier IN :entity_ids
""").bindparams( # noqa: S608 - self._tablename is not untrusted
sqlalchemy.bindparam(
key="entity_ids", value=list(uncached_ids), expanding=True
)
)

try:
with self.engine.begin() as connectable:
cur = connectable.execute(query)
except SQLAlchemyError as error:
msg = f"Unable to fetch entities by identifiers from sample store {self._tablename}"
self.log.critical(f"{msg}. Error: {error}")
raise SystemError(f"{msg}. Error: {error}") from error

# Build result dictionary to handle multiple measurement results per entity
entities_dict: dict[str, Entity] = {}
for entity_identifier, entity_representation, result_data in cur:
if entity_identifier not in entities_dict:
try:
entities_dict[entity_identifier] = Entity.model_validate(
json.loads(entity_representation)
)
# Update cache if it exists
if self._entities is not None:
self._entities[entity_identifier] = entities_dict[
entity_identifier
]
except Exception as error:
raise FailedToDecodeStoredEntityError(
entity_identifier=entity_identifier,
entity_representation=entity_representation,
cause=error,
) from error

if result_data is None:
self.log.debug(
f"Entity {entity_identifier} had no measurements associated to it."
)
continue

try:
result_dict = json.loads(result_data)
if not result_dict.get("measurements", None):
continue

measurement_result = ValidMeasurementResult.model_validate(result_dict)
except Exception as error:
raise FailedToDecodeStoredMeasurementResultForEntityError(
entity_identifier=entity_identifier,
result_representation=result_data,
cause=error,
) from error

# Add measurement result to entity
entities_dict[entity_identifier].add_measurement_result(
result=measurement_result
)

return cached_entities + list(entities_dict.values())

def get_entities(
self,
Expand Down Expand Up @@ -932,10 +856,12 @@ def entities_in_operations(self, operation_ids: str | set[str]) -> list[Entity]:
Returns:
List of Entity objects that were sampled in the specified operation(s)
"""
# Use entity_identifiers_in_operations + entities_with_identifiers so that
# Use entity_identifiers_in_operations + get_entities so that
# the entity cache is used when fetching entities.
entity_ids = self.entity_identifiers_in_operations(operation_ids)
return self.entities_with_identifiers(entity_ids)
return self.get_entities(
identifiers=set(entity_ids), require_measurements=False
)

def entities_in_operation(self, operation_ids: str | set[str]) -> list[Entity]:
"""Deprecated: use entities_in_operations instead."""
Expand Down Expand Up @@ -1164,67 +1090,19 @@ def delete(self) -> None:
pass

def entityWithIdentifier(self, entityIdentifier: str) -> Entity | None:
"""Returns entity if its in receiver otherwise returns None"""
"""Deprecated: use :meth:`get_entities` instead.

query = sqlalchemy.text(f"""
SELECT ent.identifier, ent.representation, res.data
FROM (
SELECT identifier, representation
FROM {self._tablename} ent
WHERE identifier = :identifier
) ent
LEFT OUTER JOIN {self._tablename}_measurement_results res ON ent.identifier = res.entity_id
""").bindparams( # noqa: S608 - self._tablename is not untrusted
identifier=entityIdentifier
Returns entity if it is in the store, otherwise returns ``None``.
"""
warnings.warn(
"entityWithIdentifier is deprecated, use get_entities instead.",
DeprecationWarning,
stacklevel=2,
)

try:
with self.engine.begin() as connectable:
cur = connectable.execute(query)
except SQLAlchemyError as error:
msg = f"Unable to fetch entity {entityIdentifier} and measurements from sample store {self._tablename}"
self.log.critical(f"{msg}. Error: {error}")
raise SystemError(f"{msg}. Error: {error}") from error

entity = None
failures = 0
for entity_identifier, entity_representation, result_data in cur:
if entity is None:
try:
entity = Entity.model_validate(json.loads(entity_representation))
except Exception as error:
self.log.warning(
f"Unable to decode representation for entity {entity_identifier}.\n"
f"Representation was: {entity_representation}.\n"
f"Error was {error}"
)
return None

if result_data is None:
self.log.info(
f"Entity {entity_identifier} had no measurements associated to it."
)
continue

try:
result_dict = json.loads(result_data)
if not result_dict.get("measurements", None):
continue

measurement_result = ValidMeasurementResult.model_validate(result_dict)
except Exception as error:
self.log.warning(
f"Unable to decode a measurement result for entity {entity_identifier}.\n"
f"Data was: {result_data}.\n"
f"Error was {error}"
)
failures += 1
continue

# We need to manually add valid measurements to the entity
entity.add_measurement_result(result=measurement_result)

return entity
results = self.get_entities(
identifiers=entityIdentifier, require_measurements=True
)
return results[0] if results else None

@property
def uri(self) -> str:
Expand Down Expand Up @@ -2076,7 +1954,10 @@ def _measurement_requests_cursor_to_pydantic(

# We also need the entity referenced by the measurement
if entity_id not in self._entities:
self._entities[entity_id] = self.entityWithIdentifier(entity_id)
results = self.get_entities(
identifiers=entity_id, require_measurements=False
)
self._entities[entity_id] = results[0] if results else None

entity = self._entities[entity_id]

Expand Down
5 changes: 3 additions & 2 deletions ado/modules/operators/discovery_space_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,9 +211,10 @@ async def entitiesSlice(self, start: int = 0, stop: int = 1) -> list["Entity"]:

def storedEntityWithIdentifier(self, entityIdentifier: str) -> "Entity | None":

return self._discoverySpace.sample_store.entityWithIdentifier(
entityIdentifier=entityIdentifier
results = self._discoverySpace.sample_store.get_entities(
identifiers=entityIdentifier, require_measurements=False
)
return results[0] if results else None

def storedEntitiesWithConstitutivePropertyValues(
self, propVals: list[PropertyValue]
Expand Down
2 changes: 1 addition & 1 deletion plugins/operators/anomalous_series/VERSION
Original file line number Diff line number Diff line change
@@ -1 +1 @@
1.0.4
1.0.5
Original file line number Diff line number Diff line change
Expand Up @@ -123,7 +123,7 @@ def example_configuration(cls) -> "DetectAnomalousSeries":
""",
configuration_model=DetectAnomalousSeries,
example_configuration=DetectAnomalousSeries.example_configuration(),
version="1.0.4",
version="1.0.5",
)
def detect_anomalous_series(
discoverySpace: DiscoverySpace,
Expand Down
2 changes: 1 addition & 1 deletion plugins/operators/profile_space/VERSION
Original file line number Diff line number Diff line change
@@ -1 +1 @@
2.0.3
2.0.4
4 changes: 2 additions & 2 deletions plugins/operators/profile_space/profile_space/operator.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ class ProfileParameters(pydantic.BaseModel):
# for documentation on the decorator and its parameters
@characterize_operation(
name="profile",
version="2.0.3",
version="2.0.4",
configuration_model=ProfileParameters,
example_configuration=ProfileParameters(),
description="Returns a data_profiling ProfileReport for the space",
Expand All @@ -40,7 +40,7 @@ def profile(
df = pd.DataFrame(
data=[
e.seriesRepresentation()
for e in discoverySpace.sample_store.entities
for e in discoverySpace.sample_store.get_entities(require_measurements=True)
if len(e.observedPropertyValues) > 0
]
)
Expand Down
2 changes: 1 addition & 1 deletion plugins/operators/ray_tune/VERSION
Original file line number Diff line number Diff line change
@@ -1 +1 @@
2.0.5
2.0.6
2 changes: 1 addition & 1 deletion plugins/operators/ray_tune/ado_ray_tune/operator.py
Original file line number Diff line number Diff line change
Expand Up @@ -980,7 +980,7 @@ def operator_metadata(cls) -> OperatorMetadata:
"""Returns operator metadata for the ray_tune explore operator."""
return OperatorMetadata(
name="ray_tune",
version="2.0.5",
version="2.0.6",
description=cls.description(),
configuration_model=RayTuneConfiguration,
example_configuration=RayTuneConfiguration(
Expand Down
10 changes: 7 additions & 3 deletions plugins/operators/ray_tune/ado_ray_tune/rifferla.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,7 +93,7 @@ def example_configuration(cls) -> "RifferlaParameters":
"It does this by identifying which entity space dimensions should be fixed to set values, which explored, and setting range limits for those dimensions. "
"The method leverages Mutual Information to identify dimensions correlated with the desired observed property.",
example_configuration=RifferlaParameters.example_configuration(),
version="2.0.3",
version="2.0.6",
)
def rifferla(
discoverySpace: DiscoverySpace,
Expand Down Expand Up @@ -175,7 +175,9 @@ def rifferla(
all_entities = discoverySpace.matchingEntities()
else:
print("Getting all entities")
all_entities = discoverySpace.sample_store.entities
all_entities = discoverySpace.sample_store.get_entities(
require_measurements=True
)

print(f"Number of entities: {len(all_entities)}")

Expand Down Expand Up @@ -344,7 +346,9 @@ def rifferla(
print(
"Rifferla space refinement] ...done. Copying entities (since it is a probabilistic space)..."
)
copy_entities = list(discoverySpace.sample_store.entities)
copy_entities = discoverySpace.sample_store.get_entities(
require_measurements=True
)
new_discovery_space.sample_store.add_external_entities(copy_entities)

print(
Expand Down
5 changes: 4 additions & 1 deletion tests/samplestore/create/test_create.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,7 +177,10 @@ def test_add_external_entities(
assert len(sample_store.entity_identifiers().intersection({entity.identifier})) == 1

#
retrieved_entity = sample_store.entityWithIdentifier(entity.identifier)
results = sample_store.get_entities(
identifiers=entity.identifier, require_measurements=False
)
retrieved_entity = results[0] if results else None
assert retrieved_entity is not None
assert len(retrieved_entity.propertyValues) == len(entity.propertyValues)
for i, property_value in enumerate(entity.propertyValues):
Expand Down
Loading