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
2 changes: 1 addition & 1 deletion adbc_drivers_validation/compare.py
Original file line number Diff line number Diff line change
Expand Up @@ -300,7 +300,7 @@ def compare_tables(
expected: pyarrow.Table,
actual: pyarrow.Table,
meta: query_metadata.QueryMetadata | None = None,
):
) -> None:
"""Compare two Arrow tables for equality."""
expected = make_nullable(expected)
actual = make_nullable(actual)
Expand Down
4 changes: 2 additions & 2 deletions adbc_drivers_validation/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,7 +85,7 @@ class FromEnv(BaseModel):

env: str = Field(description="Environment variable to read the value from")

def __init__(self, env: str):
def __init__(self, env: str) -> None:
super().__init__(env=env)

def get_or_raise(self) -> str:
Expand Down Expand Up @@ -141,7 +141,7 @@ class DriverFeatures(BaseModel):
quirk_get_objects_constraints_primary_normalized: bool = Field(default=False)
quirk_get_objects_constraints_unique_normalized: bool = Field(default=False)

def __init__(self, **data: typing.Any):
def __init__(self, **data: typing.Any) -> None:
super().__init__(**data)
self._current_catalog = data.get("current_catalog")
self._current_schema = data.get("current_schema")
Expand Down
2 changes: 1 addition & 1 deletion adbc_drivers_validation/quirks.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ def split_statement(statement: str, dialect: str | None = None) -> list[str]:
statements = []
current = []

def flush():
def flush() -> None:
nonlocal statements
nonlocal current
v = "\n".join(current).strip()
Expand Down
2 changes: 1 addition & 1 deletion adbc_drivers_validation/tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@
pytest.register_assert_rewrite("adbc_drivers_validation.tests.statement")


def pytest_addoption(parser):
def pytest_addoption(parser) -> None:
parser.addoption("--repl", action="store_true", default=False)
parser.addoption("--show-queries", action="store_true", default=False)

Expand Down
4 changes: 2 additions & 2 deletions adbc_drivers_validation/tests/connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,7 @@ def test_set_current_catalog(
self,
driver: model.DriverQuirks,
conn_factory: typing.Callable[[], adbc_driver_manager.dbapi.Connection],
):
) -> None:
if not driver.features.connection_set_current_catalog:
pytest.skip("not implemented")

Expand All @@ -104,7 +104,7 @@ def test_set_current_schema(
self,
driver: model.DriverQuirks,
conn_factory: typing.Callable[[], adbc_driver_manager.dbapi.Connection],
):
) -> None:
if not driver.features.connection_set_current_schema:
pytest.skip("not implemented")

Expand Down
18 changes: 9 additions & 9 deletions tests/test_arrowjson.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
from adbc_drivers_validation import arrowjson


def test_parse_int32_schema():
def test_parse_int32_schema() -> None:
f = io.StringIO("""
{
"format": "+s",
Expand All @@ -44,7 +44,7 @@ def test_parse_int32_schema():
assert parsed_schema.field("res").nullable


def test_parse_extension_schema():
def test_parse_extension_schema() -> None:
f = io.StringIO("""
{
"format": "+s",
Expand All @@ -71,15 +71,15 @@ def test_parse_extension_schema():
}


def test_parse_type_format_primitives():
def test_parse_type_format_primitives() -> None:
assert arrowjson.parse_type_format("n") == pyarrow.null()
assert arrowjson.parse_type_format("b") == pyarrow.bool_()
assert arrowjson.parse_type_format("i") == pyarrow.int32()
assert arrowjson.parse_type_format("l") == pyarrow.int64()
assert arrowjson.parse_type_format("g") == pyarrow.float64()


def test_parse_type_format_nested():
def test_parse_type_format_nested() -> None:
# Simple struct
struct_type = arrowjson.parse_type_format(
"+s",
Expand All @@ -106,7 +106,7 @@ def test_parse_type_format_nested():
assert fixed_list_type.list_size == 3


def test_from_dict_simple():
def test_from_dict_simple() -> None:
field = arrowjson.field_from_dict(
{"name": "test_field", "format": "i", "flags": ["nullable"]}
)
Expand All @@ -116,7 +116,7 @@ def test_from_dict_simple():
assert field.nullable


def test_from_dict_nested():
def test_from_dict_nested() -> None:
field = arrowjson.field_from_dict(
{
"name": "parent",
Expand Down Expand Up @@ -145,7 +145,7 @@ def test_from_dict_nested():
assert not child2.nullable


def test_from_dict_list():
def test_from_dict_list() -> None:
field = arrowjson.field_from_dict(
{
"name": "items",
Expand All @@ -161,7 +161,7 @@ def test_from_dict_list():
assert field.type.value_type == pyarrow.float32()


def test_load_table_extra_fields():
def test_load_table_extra_fields() -> None:
schema = pyarrow.schema([pyarrow.field("a", pyarrow.int32())])
data = [{"a": 1, "b": 2}]
with pytest.raises(ValueError, match="Extra fields in row: {'b': 2}"):
Expand Down Expand Up @@ -249,7 +249,7 @@ def test_array_from_values_struct() -> None:
assert actual == expected


def test_parse_type_format_decimal():
def test_parse_type_format_decimal() -> None:
"""Test decimal format parsing with various bitwidths and legacy format."""
# Test decimal32 (precision 1-9)
decimal32_type = arrowjson.parse_type_format("d:9,2,32")
Expand Down
4 changes: 2 additions & 2 deletions tests/test_compare.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
import adbc_drivers_validation.compare as compare


def test_compare_fields():
def test_compare_fields() -> None:
f1 = pyarrow.field("a", pyarrow.int32(), nullable=True)
f2 = pyarrow.field("b", pyarrow.int32(), nullable=True)
f3 = pyarrow.field("a", pyarrow.int64(), nullable=True)
Expand Down Expand Up @@ -58,7 +58,7 @@ def test_compare_fields():
compare.compare_fields(f1, f6)


def test_compare_schemas_nullability():
def test_compare_schemas_nullability() -> None:
# Should ignore nullability
s1 = pyarrow.schema([pyarrow.field("a", pyarrow.int32(), nullable=True)])
s2 = pyarrow.schema([pyarrow.field("a", pyarrow.int32(), nullable=False)])
Expand Down
34 changes: 17 additions & 17 deletions tests/test_query_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@
ROOT = Path(__file__).parent.parent


def test_connection_options_parse():
def test_connection_options_parse() -> None:
opts = ConnectionOptions.model_validate(
{
"options": {
Expand All @@ -47,7 +47,7 @@ def test_connection_options_parse():
assert opts.options["complex"].revert == "off"


def test_statement_options_parse():
def test_statement_options_parse() -> None:
opts = StatementOptions.model_validate(
{
"options": {
Expand All @@ -65,7 +65,7 @@ def test_statement_options_parse():
class TestTagsMetadata:
"""Test the TagsMetadata model."""

def test_empty_tags(self):
def test_empty_tags(self) -> None:
"""Test creating empty tags."""
tags = TagsMetadata()
assert tags.sql_type_name is None
Expand All @@ -74,17 +74,17 @@ def test_empty_tags(self):
assert tags.broken_driver is None
assert tags.show_arrow_type_parameters is False

def test_tags_with_sql_type_name(self):
def test_tags_with_sql_type_name(self) -> None:
"""Test tags with sql-type-name."""
tags = TagsMetadata.model_validate({"sql-type-name": "VARCHAR"})
assert tags.sql_type_name == "VARCHAR"

def test_tags_with_caveats(self):
def test_tags_with_caveats(self) -> None:
"""Test tags with caveats."""
tags = TagsMetadata(caveats=["Note 1", "Note 2"])
assert tags.caveats == ["Note 1", "Note 2"]

def test_tags_with_partial_support(self):
def test_tags_with_partial_support(self) -> None:
"""Test tags with partial-support."""
tags = TagsMetadata.model_validate({"partial-support": True})
assert tags.partial_support is True
Expand All @@ -93,19 +93,19 @@ def test_tags_with_partial_support(self):
class TestSetupMetadata:
"""Test the SetupMetadata model."""

def test_empty_setup(self):
def test_empty_setup(self) -> None:
"""Test creating empty setup."""
setup = SetupMetadata()
assert setup.drop is None
assert setup.connection is None
assert setup.statement is None

def test_setup_with_drop(self):
def test_setup_with_drop(self) -> None:
"""Test setup with drop table."""
setup = SetupMetadata(drop="test_table")
assert setup.drop == "test_table"

def test_setup_with_connection(self):
def test_setup_with_connection(self) -> None:
"""Test setup with connection options."""
setup = SetupMetadata(
connection=ConnectionOptions.model_validate({"options": {"key": "value"}})
Expand All @@ -118,7 +118,7 @@ def test_setup_with_connection(self):
class TestQueryMetadata:
"""Test the QueryMetadata model."""

def test_empty_metadata(self):
def test_empty_metadata(self) -> None:
"""Test creating empty metadata."""
metadata = QueryMetadata()
assert metadata.hide is False
Expand All @@ -129,30 +129,30 @@ def test_empty_metadata(self):
assert metadata.statement is None
assert isinstance(metadata.tags, TagsMetadata)

def test_metadata_with_hide(self):
def test_metadata_with_hide(self) -> None:
"""Test metadata with hide flag."""
metadata = QueryMetadata(hide=True)
assert metadata.hide is True

def test_metadata_with_skip(self):
def test_metadata_with_skip(self) -> None:
"""Test metadata with skip reason."""
metadata = QueryMetadata(skip="Not supported")
assert metadata.skip == "Not supported"

def test_metadata_with_sort_keys(self):
def test_metadata_with_sort_keys(self) -> None:
"""Test metadata with sort-keys."""
metadata = QueryMetadata.model_validate(
{"sort-keys": [("col1", "ascending"), ("col2", "descending")]}
)
assert metadata.sort_keys == [("col1", "ascending"), ("col2", "descending")]

def test_metadata_with_setup(self):
def test_metadata_with_setup(self) -> None:
"""Test metadata with setup section."""
metadata = QueryMetadata(setup=SetupMetadata(drop="test_table"))
assert metadata.setup is not None
assert metadata.setup.drop == "test_table"

def test_metadata_with_tags(self):
def test_metadata_with_tags(self) -> None:
"""Test metadata with tags."""
metadata = QueryMetadata(
tags=TagsMetadata.model_validate({"sql-type-name": "INTEGER"})
Expand All @@ -170,7 +170,7 @@ def test_metadata_with_tags(self):
for path in _root.rglob("*.toml")
],
)
def test_load_all_toml(path: Path):
def test_load_all_toml(path: Path) -> None:
with path.open("rb") as f:
data = tomllib.load(f)
QueryMetadata.model_validate(data)
Expand All @@ -183,7 +183,7 @@ def test_load_all_toml(path: Path):
for path in _root.rglob("*.txtcase")
],
)
def test_load_all_txtcase(path: Path):
def test_load_all_txtcase(path: Path) -> None:
t = adbc_drivers_validation.txtcase.load(path)
try:
data = t.get_part("metadata")
Expand Down
8 changes: 4 additions & 4 deletions tests/test_queryset.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,15 +20,15 @@
from adbc_drivers_validation import model, query_metadata


def test_txtcase_empty(tmp_path: Path):
def test_txtcase_empty(tmp_path: Path) -> None:
with (tmp_path / "query.txtcase").open("w") as f:
f.write("\n")

with pytest.raises(ValueError, match="unknown query type"):
model.QuerySet.load(tmp_path)


def test_txtcase_select(tmp_path: Path):
def test_txtcase_select(tmp_path: Path) -> None:
with (tmp_path / "query.txtcase").open("w") as f:
f.write(
"""
Expand Down Expand Up @@ -70,7 +70,7 @@ def test_txtcase_select(tmp_path: Path):
assert query.query.expected_result().to_pylist() == [{"$0": 1}]


def test_txtcase_override(tmp_path: Path):
def test_txtcase_override(tmp_path: Path) -> None:
(tmp_path / "base").mkdir()
(tmp_path / "over").mkdir()

Expand Down Expand Up @@ -135,5 +135,5 @@ def test_txtcase_override(tmp_path: Path):
assert query.query.expected_result().to_pylist() == [{"$0": 2}]


def test_load_queries():
def test_load_queries() -> None:
model.base_query_set()
Loading
Loading