Skip to content
Open
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
32 changes: 22 additions & 10 deletions featuretools/entityset/entityset.py
Original file line number Diff line number Diff line change
Expand Up @@ -1470,8 +1470,9 @@ def query_by_values(

Args:
dataframe_name (str): The id of the dataframe to query
instance_vals (pd.Dataframe, pd.Series, list[str] or str) :
Instance(s) to match.
instance_vals (None, pd.Series, or iterable) :
Instance(s) to match. Values must be provided as a
``pd.Series`` or an iterable of instance values.
Comment on lines +1473 to +1475
column_name (str) : Column to query on. If None, query on index.
columns (list[str]) : Columns to return. Return all columns if None.
time_last (pd.TimeStamp) : Query data up to and including this
Expand Down Expand Up @@ -1681,21 +1682,32 @@ def replace(x):

def _vals_to_series(instance_vals, column_id):
"""
instance_vals may be a pd.Dataframe, a pd.Series, a list, a single
value, or None. This function always returns a Series or None.
instance_vals may be None, a pd.Series, or an iterable of values.
This function always returns a Series or None.
Comment on lines +1685 to +1686
"""
if instance_vals is None:
return None

# If this is a single value, make it a list
if not hasattr(instance_vals, "__iter__"):
instance_vals = [instance_vals]
if isinstance(instance_vals, str):
raise TypeError(
"instance_vals must be a pd.Series or an iterable of values, "
"not a string.",
)

# convert iterable to pd.Series
if isinstance(instance_vals, pd.DataFrame):
out_vals = instance_vals[column_id]
else:
raise TypeError(
"instance_vals must be a pd.Series or an iterable of values, "
"not a pd.DataFrame.",
)

if isinstance(instance_vals, pd.Series):
out_vals = instance_vals
elif hasattr(instance_vals, "__iter__"):
out_vals = pd.Series(instance_vals)
else:
raise TypeError(
"instance_vals must be a pd.Series or an iterable of values.",
)

# no duplicates or NaN values
out_vals = out_vals.drop_duplicates().dropna()
Expand Down
21 changes: 15 additions & 6 deletions featuretools/tests/entityset_tests/test_es.py
Original file line number Diff line number Diff line change
Expand Up @@ -346,16 +346,25 @@ def test_query_by_id(es):
assert df["id"].values[0] == 0


def test_query_by_single_value(es):
df = es.query_by_values("log", instance_vals=0)
def test_query_by_series(es):
df = es.query_by_values("log", instance_vals=pd.Series([0]))
assert df["id"].values[0] == 0


def test_query_by_df(es):
instance_df = pd.DataFrame({"id": [1, 3], "vals": [0, 1]})
df = es.query_by_values("log", instance_vals=instance_df)
def test_query_by_single_value_rejected(es):
with pytest.raises(TypeError, match="instance_vals must be a pd.Series"):
es.query_by_values("log", instance_vals=0)


def test_query_by_string_rejected(es):
with pytest.raises(TypeError, match="instance_vals must be a pd.Series"):
es.query_by_values("log", instance_vals="0")

assert np.array_equal(df["id"], [1, 3])

def test_query_by_df_rejected(es):
instance_df = pd.DataFrame({"id": [1, 3], "vals": [0, 1]})
with pytest.raises(TypeError, match="instance_vals must be a pd.Series"):
es.query_by_values("log", instance_vals=instance_df)


def test_query_by_id_with_time(es):
Expand Down