From 85f0e98a3fef3de07353f39c2ca6dda48fcfb2e9 Mon Sep 17 00:00:00 2001 From: Michael Rogozin Date: Tue, 1 Oct 2024 01:28:10 +0300 Subject: [PATCH 1/3] fix: resolve tests imports problem --- setup.py | 2 +- tests/__init__.py | 0 tests/test_base.py | 110 +++++++++--------- tests/test_crud.py | 74 +++++++----- tests/test_dialect.py | 263 +++++++++++++++++++++--------------------- tests/test_orm_dao.py | 61 +++++----- tests/test_orm_dto.py | 70 ++++++----- 7 files changed, 306 insertions(+), 274 deletions(-) create mode 100644 tests/__init__.py diff --git a/setup.py b/setup.py index fae6944..073eb0e 100644 --- a/setup.py +++ b/setup.py @@ -22,7 +22,7 @@ python_requires = '>=3.9', install_requires = ['sqlalchemy==2.0.27', - 'pysqream==3.2.5', + 'pysqream>=3.2.5', 'setuptools>=57.4.0', 'pandas==2.2.1', 'numpy>=1.20', diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/test_base.py b/tests/test_base.py index e2c4819..cb867ab 100644 --- a/tests/test_base.py +++ b/tests/test_base.py @@ -1,33 +1,37 @@ import socket -from datetime import datetime, date +from datetime import date, datetime from decimal import Decimal -from random import choice, randint, choices +from random import choice, choices, randint from typing import Union import pytest import sqlalchemy as sa -from sqlalchemy import (text, - Table, - Column, - orm, - Integer, - Boolean, - Date, - DateTime, - Numeric, - Text, - create_engine, - MetaData, - Identity, - Connection) -from sqlalchemy.orm import declarative_base, Session - -from pytest_logger import Logger +from sqlalchemy import ( + Boolean, + Column, + Connection, + Date, + DateTime, + Identity, + Integer, + MetaData, + Numeric, + Table, + Text, + create_engine, + orm, + text, +) +from sqlalchemy.dialects import registry +from sqlalchemy.orm import Session, declarative_base + +from tests.pytest_logger import Logger def connect(ip, port, clustered=False, use_ssl=False): print_echo = False conn_str = f"pysqream+dialect://sqream:sqream@{ip}:{port}/master" + registry.register("pysqream.dialect", "pysqream_sqlalchemy.dialect", "SqreamDialect") engine = create_engine(conn_str, echo=print_echo, connect_args={"clustered": clustered, "use_ssl": use_ssl}) sa.Tinyint = engine.dialect.Tinyint session = orm.sessionmaker(bind=engine)() @@ -41,11 +45,11 @@ def setTinyint(engine): class TestBase: - @pytest.fixture() + @pytest.fixture def ip(self, pytestconfig): return pytestconfig.getoption("ip") - @pytest.fixture() + @pytest.fixture def port(self, pytestconfig): return pytestconfig.getoption("port") @@ -71,7 +75,7 @@ def stop(self): class TestBaseOrm(TestBase): - @pytest.fixture() + @pytest.fixture def ip(self, pytestconfig): return pytestconfig.getoption("ip") @@ -79,27 +83,27 @@ def ip(self, pytestconfig): def port(self, pytestconfig): return pytestconfig.getoption("port") - @pytest.fixture() + @pytest.fixture def Base(self, pytestconfig): return self.Base - @pytest.fixture() + @pytest.fixture def user(self): return self.user - @pytest.fixture() + @pytest.fixture def address(self): return self.address - @pytest.fixture() + @pytest.fixture def dates(self): return self.dates - @pytest.fixture() + @pytest.fixture def table1(self): return self.table1 - @pytest.fixture() + @pytest.fixture def table2(self): return self.table2 @@ -143,13 +147,13 @@ def __repr__(self): self.start(ip, port) self.table1 = Table( - 'table1', self.metadata, - Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table1", self.metadata, + Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) self.table2 = Table( - 'table2', self.metadata, - Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table2", self.metadata, + Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) yield @@ -157,11 +161,11 @@ def __repr__(self): class TestBaseTI(TestBase): - @pytest.fixture() + @pytest.fixture def ip(self, pytestconfig): return pytestconfig.getoption("ip") - @pytest.fixture() + @pytest.fixture def testware_affinity_matrix(self): return self.testware_affinity_matrix @@ -170,17 +174,17 @@ def Test_setup_teardown(self, ip, port): self.start(ip, port) self.testware_affinity_matrix = sa.Table( - 'testware_affinity_matrix', + "testware_affinity_matrix", self.metadata, - sa.Column('technology', sa.TEXT(32)), - sa.Column('criteria', sa.TEXT(32)), - sa.Column('category', sa.TEXT(32)), - sa.Column('component', sa.TEXT(32)), - sa.Column('svn', sa.TEXT(32)), - sa.Column('parm_name', sa.TEXT(32)), - sa.Column('lpt', sa.TEXT(32)), - sa.Column('tech', sa.TEXT(32)), - sa.Column('severity', sa.Float) + sa.Column("technology", sa.TEXT(32)), + sa.Column("criteria", sa.TEXT(32)), + sa.Column("category", sa.TEXT(32)), + sa.Column("component", sa.TEXT(32)), + sa.Column("svn", sa.TEXT(32)), + sa.Column("parm_name", sa.TEXT(32)), + sa.Column("lpt", sa.TEXT(32)), + sa.Column("tech", sa.TEXT(32)), + sa.Column("severity", sa.Float), ) if self.insp.has_table(self.testware_affinity_matrix.name): @@ -193,7 +197,7 @@ def Test_setup_teardown(self, ip, port): class TestBaseCRUD(TestBase): - database_name = schema_name = table_name = 'crud' + database_name = schema_name = table_name = "crud" view_name = "view_for_crud" @staticmethod @@ -228,19 +232,19 @@ def crud_table(self): return Table( self.table_name, self.metadata, - Column('i', Integer), - Column('b', Boolean), - Column('d', Date), - Column('dt', DateTime), - Column('n', Numeric(15, 6)), - Column('t', Text), + Column("i", Integer), + Column("b", Boolean), + Column("d", Date), + Column("dt", DateTime), + Column("n", Numeric(15, 6)), + Column("t", Text), # Column('iar', ARRAY(Integer)), # Column('bar', ARRAY(Boolean)), # Column('dar', ARRAY(Date)), # Column('dtar',ARRAY(DateTime)), # Column('nar', ARRAY(Numeric(15, 6))), # Column('tar', ARRAY(Text)), - extend_existing=True + extend_existing=True, ) @staticmethod @@ -256,7 +260,7 @@ def get_random_row_values_for_crud_table(row_number: int): minute=randint(1, 59), second=randint(1, 59)), Decimal(f"{randint(int(1e8), int(9e8))}.{randint(int(1e5), int(9e5))}"), - "".join(choices("ABCDEFGHIJKLMNOPQRSTUVWXYZ", k=randint(5, 50))) + "".join(choices("ABCDEFGHIJKLMNOPQRSTUVWXYZ", k=randint(5, 50))), ) def recreate_all_via_metadata(self, executor: Union[Connection, Session] = None): diff --git a/tests/test_crud.py b/tests/test_crud.py index b8c6ded..0ee17b6 100644 --- a/tests/test_crud.py +++ b/tests/test_crud.py @@ -1,19 +1,16 @@ -import sys -sys.path.insert(0, 'pysqream_sqlalchemy') -sys.path.insert(0, 'tests') -from datetime import datetime, date + +from datetime import date, datetime from decimal import Decimal from typing import Union import pytest -from sqlalchemy import text, Table, Column, dialects, Integer, Text, select, Connection, insert -from sqlalchemy.orm import sessionmaker, Session +from sqlalchemy import Column, Connection, Integer, Table, Text, dialects, select, text, insert from sqlalchemy.engine.row import Row +from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.sql import text +from tests.test_base import TestBaseCRUD -from test_base import TestBaseCRUD - - -dialects.registry.register("pysqream.dialect", "dialect", "SqreamDialect") +dialects.registry.register("pysqream.dialect", "pysqream_sqlalchemy.dialect", "SqreamDialect") class TestCreate(TestBaseCRUD): @@ -218,7 +215,7 @@ def insert_values_into_table(self, ( ("select(crud_table)", 100), ("select(crud_table).where(text('i>50'))", 50), - ) + ), ) def test_read_from_table(self, crud_table, statement, expected_result): self.recreate_all_via_metadata() @@ -231,12 +228,28 @@ def test_read_from_table(self, crud_table, statement, expected_result): assert len(result1) == expected_result == len(result2) + # def test_read_from_table_parametrized(self, crud_table): + # self.recreate_all_via_metadata() + # + # self.insert_values_into_table(crud_table, rows_amount=10) + # + # # Write the select query using SQLAlchemy + # query = select([crud_table.c.i]).where(crud_table.c.i == text(":value")) + # + # # Execute the query with parameter binding + # result = connection.execute(query, value=123) # Replace 123 with actual value + # + # # Fetch and print results + # for row in result: + # print(row['i']) + + @pytest.mark.parametrize( ("statement", "expected_result"), ( ("select(crud_table)", 100), ("select(crud_table).where(text('i>50'))", 50), - ) + ), ) def test_read_from_table_with_engine_context_manager(self, crud_table, statement, expected_result): with self.engine.connect() as connection: @@ -255,7 +268,7 @@ def test_read_from_table_with_engine_context_manager(self, crud_table, statement ( ("select(crud_table)", 100), ("select(crud_table).where(text('i>50'))", 50), - ) + ), ) def test_read_from_table_with_session_context_manager(self, crud_table, statement, expected_result): new_session = sessionmaker(self.engine) @@ -313,14 +326,14 @@ def test_read_from_view_with_session_context_manager(self, crud_table): assert len(result1) == 10 == len(result2) @pytest.mark.parametrize( - ('query', "entity_name"), + ("query", "entity_name"), ( ("select * from sqream_catalog.tables where table_name = '{}'", TestBaseCRUD.table_name), ("select * from sqream_catalog.views where view_name = '{}'", TestBaseCRUD.view_name), ("select * from sqream_catalog.chunks join sqream_catalog.tables on " "sqream_catalog.tables.table_id = sqream_catalog.chunks.table_id where table_name = '{}'", TestBaseCRUD.table_name), - ) + ), ) def test_read_from_sqream_catalog(self, crud_table, query, entity_name): self.recreate_all_via_metadata() @@ -338,14 +351,14 @@ def test_read_from_sqream_catalog(self, crud_table, query, entity_name): assert len(result1) == 1 == len(result2) @pytest.mark.parametrize( - ('query', "entity_name"), + ("query", "entity_name"), ( ("select * from sqream_catalog.tables where table_name = '{}'", TestBaseCRUD.table_name), ("select * from sqream_catalog.views where view_name = '{}'", TestBaseCRUD.view_name), ("select * from sqream_catalog.chunks join sqream_catalog.tables on " "sqream_catalog.tables.table_id = sqream_catalog.chunks.table_id where table_name = '{}'", TestBaseCRUD.table_name), - ) + ), ) def test_read_from_sqream_catalog_with_engine_context_manager(self, crud_table, query, entity_name): with self.engine.connect() as connection: @@ -364,14 +377,14 @@ def test_read_from_sqream_catalog_with_engine_context_manager(self, crud_table, assert len(result1) == 1 == len(result2) @pytest.mark.parametrize( - ('query', "entity_name"), + ("query", "entity_name"), ( ("select * from sqream_catalog.tables where table_name = '{}'", TestBaseCRUD.table_name), ("select * from sqream_catalog.views where view_name = '{}'", TestBaseCRUD.view_name), ("select * from sqream_catalog.chunks join sqream_catalog.tables on " "sqream_catalog.tables.table_id = sqream_catalog.chunks.table_id where table_name = '{}'", TestBaseCRUD.table_name), - ) + ), ) def test_read_from_sqream_catalog_with_session_context_manager(self, crud_table, query, entity_name): new_session = sessionmaker(self.engine) @@ -392,7 +405,7 @@ def test_read_from_sqream_catalog_with_session_context_manager(self, crud_table, def test_read_with_join(self, crud_table): table1 = Table( - 'table1', + "table1", self.metadata, Column("i", Integer), Column("t", Text), @@ -400,7 +413,7 @@ def test_read_with_join(self, crud_table): self.recreate_all_via_metadata() - values1 = [(i, 't' * i) for i in range(1, 11)] + values1 = [(i, "t" * i) for i in range(1, 11)] values2 = [self.get_random_row_values_for_crud_table(i) for i in range(1, 11)] self.insert_values_into_table(crud_table, rows_amount=10, values=values2) self.insert_values_into_table(table1, rows_amount=10, values=values1) @@ -417,7 +430,7 @@ def test_read_with_join(self, crud_table): def test_read_with_join_with_engine_context_manager(self, crud_table): table1 = Table( - 'table1', + "table1", self.metadata, Column("i", Integer), Column("t", Text), @@ -425,7 +438,7 @@ def test_read_with_join_with_engine_context_manager(self, crud_table): with self.engine.connect() as connection: self.recreate_all_via_metadata(executor=connection) - values1 = [(i, 't' * i) for i in range(1, 11)] + values1 = [(i, "t" * i) for i in range(1, 11)] values2 = [self.get_random_row_values_for_crud_table(i) for i in range(1, 11)] self.insert_values_into_table(crud_table, rows_amount=10, values=values2, executor=connection) self.insert_values_into_table(table1, rows_amount=10, values=values1, executor=connection) @@ -443,7 +456,7 @@ def test_read_with_join_with_engine_context_manager(self, crud_table): def test_read_with_join_with_session_context_manager(self, crud_table): new_session = sessionmaker(self.engine) table1 = Table( - 'table1', + "table1", self.metadata, Column("i", Integer), Column("t", Text), @@ -451,7 +464,7 @@ def test_read_with_join_with_session_context_manager(self, crud_table): with new_session.begin() as session: self.recreate_all_via_metadata(executor=session) - values1 = [(i, 't' * i) for i in range(1, 11)] + values1 = [(i, "t" * i) for i in range(1, 11)] values2 = [self.get_random_row_values_for_crud_table(i) for i in range(1, 11)] self.insert_values_into_table(crud_table, rows_amount=10, values=values2, executor=session) self.insert_values_into_table(table1, rows_amount=10, values=values1, executor=session) @@ -531,7 +544,7 @@ def test_session_add_delete(self, crud_table_row): ( "crud_table.insert().values(values)", "insert(crud_table).values(values)", - ) + ), ) def test_insert_into_table(self, crud_table, insert_statement): self.recreate_all_via_metadata() @@ -551,7 +564,7 @@ def test_insert_into_table(self, crud_table, insert_statement): ( "crud_table.insert().values(values)", "insert(crud_table).values(values)", - ) + ), ) def test_insert_into_table_with_engine_context_manager(self, crud_table, insert_statement): with (self.engine.connect() as connection): @@ -572,7 +585,7 @@ def test_insert_into_table_with_engine_context_manager(self, crud_table, insert_ ( "crud_table.insert().values(values)", "insert(crud_table).values(values)", - ) + ), ) def test_insert_into_table_with_session_context_manager(self, crud_table, insert_statement): new_session = sessionmaker(self.engine) @@ -783,8 +796,7 @@ def test_delete_view_with_session_context_manager(self, crud_table): class TestUtilityFunctions(TestBaseCRUD): - """ - I took all utility functions below from: + """I took all utility functions below from: https://docs.sqream.com/en/latest/search.html?q=utility&check_keywords=yes&area=default """ @@ -798,7 +810,7 @@ class TestUtilityFunctions(TestBaseCRUD): f"select get_data_metrics('daily', '{datetime.now()}', '{datetime.now()}')", "select get_gpu_info()", "select show_last_node_info()", - ) + ), ) def test_most_used_utility_functions(self, query): result = self.session.execute(text(query)).fetchall() diff --git a/tests/test_dialect.py b/tests/test_dialect.py index f62a410..c80062e 100644 --- a/tests/test_dialect.py +++ b/tests/test_dialect.py @@ -1,30 +1,27 @@ -""" - Testing the SQream SQLAlchemy dialect. See also tests for the SQream - DB-API connector +"""Testing the SQream SQLAlchemy dialect. See also tests for the SQream +DB-API connector """ import os -import sys +from datetime import date, datetime +from decimal import Decimal -sys.path.insert(0, 'pysqream_sqlalchemy') -sys.path.insert(0, 'tests') -import pytest import pandas as pd +import pytest import sqlalchemy as sa from sqlalchemy import create_engine, select, Table, Column, insert, text, DDL, orm, Identity, BigInteger, MetaData, engine from test_base import TestBase, Logger, TestBaseTI from alembic.runtime.migration import MigrationContext -from alembic.operations import Operations -from datetime import datetime, date -from decimal import Decimal +from sqlalchemy import DDL, Column, Table, create_engine, insert, orm, select, text + +from tests.test_base import Logger, TestBase, TestBaseTI # Registering dialect sa.dialects.registry.register("pysqream.dialect", "dialect", "SqreamDialect") def find_diff(df1: pd.DataFrame, df2: pd.DataFrame): - """ - Find the differance between two dataframes + """Find the differance between two dataframes """ if len(df1.index) != len(df2.index): msg = f"Row count does not match\nSQream returned {len(df1.index)}\nremote returned: {len(df2.index)}" @@ -36,18 +33,18 @@ def find_diff(df1: pd.DataFrame, df2: pd.DataFrame): class TestSqlalchemy(TestBase): def test_sqlalchemy(self): - Logger().info('SQLAlchemy direct query tests') + Logger().info("SQLAlchemy direct query tests") # Test 0 - as bestowed upon me by Yuval. Using the URL object directly instead of a connection string - sa.dialects.registry.register("pysqream.dialect", "dialect", "SqreamDialect") - manual_conn_str = sa.engine.url.URL.create(drivername='pysqream+dialect', - username='sqream', - password='sqream', - host=f'{self.ip}', + sa.dialects.registry.register("pysqream.dialect", "pysqream_sqlalchemy.dialect", "SqreamDialect") + manual_conn_str = sa.engine.url.URL.create(drivername="pysqream+dialect", + username="sqream", + password="sqream", + host=f"{self.ip}", port=self.port, - database='master') + database="master") engine2 = create_engine(manual_conn_str) session2 = orm.sessionmaker(bind=engine2)() - res = session2.execute(DDL('select 1')) + res = session2.execute(DDL("select 1")) assert (all(row[0] == 1 for row in res)) # Simple direct Engine query - this passes the queries to the underlying DB-API @@ -61,32 +58,32 @@ def test_sqlalchemy(self): assert (res.fetchall() == [(4,), (5,)]) # Reflection test - inspected_cols = self.insp.get_columns('kOko') - assert (inspected_cols[0]['name'] == 'iNts fosho') + inspected_cols = self.insp.get_columns("kOko") + assert (inspected_cols[0]["name"] == "iNts fosho") - self.metadata.reflect(bind=self.engine, only={'kOko'}) + self.metadata.reflect(bind=self.engine, only={"kOko"}) assert (repr(self.metadata.tables["kOko"]) == "Table('kOko', MetaData(), Column('iNts fosho', Integer(), table=, nullable=False), schema=None)") - Logger().info('SQLAlchemy ORM tests') + Logger().info("SQLAlchemy ORM tests") # ORM queries - test that correct SQream queries (SQL text strings) are # created (that are then passed to the DB-API) # Create table via ORM orm_table = Table( - 'orm_table', self.metadata, - Column('bools', sa.Boolean), - Column('ubytes', sa.Tinyint), - Column('shorts', sa.SmallInteger), - Column('iNts', sa.Integer), - Column('bigints', sa.BigInteger), - Column('floats', sa.REAL), - Column('doubles', sa.Float), - Column('dates', sa.Date), - Column('datetimes', sa.DateTime), - Column('varchars', sa.String(10)), - Column('nvarchars', sa.UnicodeText), - Column('numerics', sa.Numeric(38, 1)), - extend_existing=True + "orm_table", self.metadata, + Column("bools", sa.Boolean), + Column("ubytes", sa.Tinyint), + Column("shorts", sa.SmallInteger), + Column("iNts", sa.Integer), + Column("bigints", sa.BigInteger), + Column("floats", sa.REAL), + Column("doubles", sa.Float), + Column("dates", sa.Date), + Column("datetimes", sa.DateTime), + Column("varchars", sa.String(10)), + Column("nvarchars", sa.UnicodeText), + Column("numerics", sa.Numeric(38, 1)), + extend_existing=True, ) if self.insp.has_table(orm_table.name): orm_table.drop(bind=self.engine) @@ -95,7 +92,7 @@ def test_sqlalchemy(self): # Insert into table values = [(True, 77, 777, 7777, 77777, 7.0, 7.77777777, date(2012, 11, 23), datetime(2012, 11, 23, 16, 34, 56), - 'test', 'test_text', Decimal('7.7')), ] * 2 + "test", "test_text", Decimal("7.7")) ] * 2 stmt = orm_table.insert().values(values) self.session.execute(stmt) @@ -116,42 +113,42 @@ def test_sqlalchemy(self): class TestPandas(TestBase): def test_pandas(self): # Creating a SQream table from a Pandas DataFrame - Logger().info('Pandas tests') + Logger().info("Pandas tests") df = pd.DataFrame({ - 'bools': [True, False], - 'ubytes': [10, 11], - 'shorts': [110, 111], - 'ints': [1110, 1111], - 'bigints': [1111110, 11111111], - 'floats': [10.0, 11.0], - 'doubles': [10.1111111, 11.1111111], - 'dates': [date(2012, 11, 23), date(2012, 11, 23)], - 'datetimes': [datetime(2012, 11, 23, 16, 34, 56), datetime(2012, 11, 23, 16, 34, 56)], - 'varchars': ['koko', 'koko2'], - 'nvarchars': ['shoko', 'shoko2'], - 'numerics': [Decimal("1.1"), Decimal("-1.1")] + "bools": [True, False], + "ubytes": [10, 11], + "shorts": [110, 111], + "ints": [1110, 1111], + "bigints": [1111110, 11111111], + "floats": [10.0, 11.0], + "doubles": [10.1111111, 11.1111111], + "dates": [date(2012, 11, 23), date(2012, 11, 23)], + "datetimes": [datetime(2012, 11, 23, 16, 34, 56), datetime(2012, 11, 23, 16, 34, 56)], + "varchars": ["koko", "koko2"], + "nvarchars": ["shoko", "shoko2"], + "numerics": [Decimal("1.1"), Decimal("-1.1")], }) dtype = { - 'bools': sa.Boolean, - 'ubytes': sa.Tinyint, - 'shorts': sa.SmallInteger, - 'ints': sa.Integer, - 'bigints': sa.BigInteger, - 'floats': sa.REAL, - 'doubles': sa.Float, - 'dates': sa.Date, - 'datetimes': sa.DateTime, - 'varchars': sa.String(10), - 'nvarchars': sa.UnicodeText, - 'numerics': sa.Numeric(38, 10) + "bools": sa.Boolean, + "ubytes": sa.Tinyint, + "shorts": sa.SmallInteger, + "ints": sa.Integer, + "bigints": sa.BigInteger, + "floats": sa.REAL, + "doubles": sa.Float, + "dates": sa.Date, + "datetimes": sa.DateTime, + "varchars": sa.String(10), + "nvarchars": sa.UnicodeText, + "numerics": sa.Numeric(38, 10), } # Drop, create and insert - df.to_sql('kOko3', self.engine, if_exists='replace', index=False, dtype=dtype) + df.to_sql("kOko3", self.engine, if_exists="replace", index=False, dtype=dtype) res = pd.read_sql('select * from "kOko3"', self.engine) - res2 = pd.read_sql_table('kOko3', self.engine) + res2 = pd.read_sql_table("kOko3", self.engine) assert ((res == df).eq(True).all().iloc[0]) assert ((res2 == df).eq(True).all().iloc[0]) @@ -160,76 +157,76 @@ def test_pandas(self): # Alembic tests class TestAlembic(TestBase): def test_alembic(self): - Logger().info('Alembic tests') + Logger().info("Alembic tests") session = self.session.connection() ctx = MigrationContext.configure(session) op = Operations(ctx) try: - op.drop_table('waste') + op.drop_table("waste") except Exception: pass - t = op.create_table('waste', - Column('bools', sa.Boolean), - Column('ubytes', sa.Tinyint), - Column('shorts', sa.SmallInteger), - Column('ints', sa.Integer), - Column('bigints', sa.BigInteger), - Column('floats', sa.REAL), - Column('doubles', sa.Float), - Column('dates', sa.Date), - Column('datetimes', sa.DateTime), - Column('varchars', sa.String(10)), - Column('nvarchars', sa.UnicodeText), - Column('numerics', sa.Numeric(38, 10)), + t = op.create_table("waste", + Column("bools", sa.Boolean), + Column("ubytes", sa.Tinyint), + Column("shorts", sa.SmallInteger), + Column("ints", sa.Integer), + Column("bigints", sa.BigInteger), + Column("floats", sa.REAL), + Column("doubles", sa.Float), + Column("dates", sa.Date), + Column("datetimes", sa.DateTime), + Column("varchars", sa.String(10)), + Column("nvarchars", sa.UnicodeText), + Column("numerics", sa.Numeric(38, 10)), ) data = [ { - 'bools': True, - 'ubytes': 5, - 'shorts': 55, - 'ints': 555, - 'bigints': 5555, - 'floats': 5.0, - 'doubles': 5.5555555, - 'dates': date(2012, 11, 23), - 'datetimes': datetime(2012, 11, 23, 16, 34, 56), - 'varchars': 'bla', - 'nvarchars': 'bla2', - 'numerics': Decimal("1.1") + "bools": True, + "ubytes": 5, + "shorts": 55, + "ints": 555, + "bigints": 5555, + "floats": 5.0, + "doubles": 5.5555555, + "dates": date(2012, 11, 23), + "datetimes": datetime(2012, 11, 23, 16, 34, 56), + "varchars": "bla", + "nvarchars": "bla2", + "numerics": Decimal("1.1"), }, - {'bools': False, - 'ubytes': 6, - 'shorts': 66, - 'ints': 666, - 'bigints': 6666, - 'floats': 6.0, - 'doubles': 6.6666666, - 'dates': date(2012, 11, 24), - 'datetimes': datetime(2012, 11, 24, 16, 34, 57), - 'varchars': 'bla', - 'nvarchars': 'bla2', - 'numerics': Decimal("-1.1") - } + {"bools": False, + "ubytes": 6, + "shorts": 66, + "ints": 666, + "bigints": 6666, + "floats": 6.0, + "doubles": 6.6666666, + "dates": date(2012, 11, 24), + "datetimes": datetime(2012, 11, 24, 16, 34, 57), + "varchars": "bla", + "nvarchars": "bla2", + "numerics": Decimal("-1.1"), + }, ] op.bulk_insert(t, data) - res = self.session.execute(text('select * from waste')).fetchall() + res = self.session.execute(text("select * from waste")).fetchall() assert (res == [tuple(d.values()) for d in data]) class TestTI(TestBaseTI): - @pytest.fixture() + @pytest.fixture def path(self): return f"{os.getcwd()}/tests/data/LBC9_PLV_affinity_matrix_send.csv" - @pytest.fixture() + @pytest.fixture def insert_data(self, path): - return pd.read_csv(path).to_dict('records') + return pd.read_csv(path).to_dict("records") @pytest.mark.parametrize("case", (1, 2)) def test_ti(self, path, insert_data, case): @@ -241,8 +238,8 @@ def test_ti(self, path, insert_data, case): self.session.execute(ins, insert_data) res = self.session.execute(self.testware_affinity_matrix.select()).fetchall() - res_df = pd.DataFrame(res, columns=['technology', 'criteria', 'category', 'component', - 'svn', 'parm_name', 'lpt', 'tech', 'severity']) + res_df = pd.DataFrame(res, columns=["technology", "criteria", "category", "component", + "svn", "parm_name", "lpt", "tech", "severity"]) expected_df = pd.read_csv(path) is_equal, msg_results = find_diff(expected_df, res_df) assert is_equal, msg_results @@ -251,8 +248,8 @@ def test_ti(self, path, insert_data, case): class TestNew(TestBase): def test_1(self): table1 = Table( - 'table1', self.metadata, - Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value2", sa.Integer) + "table1", self.metadata, + Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value2", sa.Integer), ) if self.insp.has_table(table1.name): table1.drop(bind=self.engine) @@ -260,7 +257,7 @@ def test_1(self): table1.create(bind=self.engine) # Insert into table - values = [(1, 'test', 2)] + values = [(1, "test", 2)] ins = insert(table1).values(values) self.session.execute(ins) @@ -270,8 +267,8 @@ def test_1(self): assert res == values, res -class TestDeclarativeBase(TestBase): - # negative +class TestDeclarativeBase(TestBase): + # negative def test_SQ_16994_negative(self): engine_url = engine.url.URL("sqream", database="master", @@ -281,25 +278,25 @@ def test_SQ_16994_negative(self): port=int(self.port), query={}) engine_obj = create_engine(engine_url, connect_args={"clustered": False, "use_ssl": False}) - + metadata_obj = MetaData(schema="public") - + class Base(orm.DeclarativeBase): metadata = metadata_obj - + class TestInsertSQlAlchemy(Base): __tablename__ = "test_insert_sqlalchemy" id: orm.Mapped[int] = orm.mapped_column( BigInteger, primary_key=True ) - + Base.metadata.drop_all(engine_obj) - with pytest.raises(Exception) as e_info: + with pytest.raises(Exception) as e_info: Base.metadata.create_all(engine_obj) - + assert "primary key constraints are not supported by SQream" in str(e_info.value) - - + + # positive # https://docs.sqlalchemy.org/en/20/faq/ormconfiguration.html#how-do-i-map-a-table-that-has-no-primary-key def test_SQ_16994_positive_1(self): @@ -311,10 +308,10 @@ def test_SQ_16994_positive_1(self): port=int(self.port), query={}) engine_obj = create_engine(engine_url, connect_args={"clustered": False, "use_ssl": False}) - + class Base(orm.DeclarativeBase): metadata = MetaData() - + class TestInsertSQlAlchemy(Base): __tablename__ = "test_insert_sqlalchemy_3" __mapper_args__ = { @@ -323,10 +320,10 @@ class TestInsertSQlAlchemy(Base): id: orm.Mapped[int] = orm.mapped_column( BigInteger ) - + Base.metadata.drop_all(engine_obj) Base.metadata.create_all(engine_obj) - + # positive def test_SQ_16994_positive_2(self): engine_url = engine.url.URL("sqream", @@ -337,15 +334,15 @@ def test_SQ_16994_positive_2(self): port=int(self.port), query={}) engine_obj = create_engine(engine_url, connect_args={"clustered": False, "use_ssl": False}) - + class Base(orm.DeclarativeBase): metadata = MetaData() - + class TestInsertSQlAlchemy(Base): __tablename__ = "test_insert_sqlalchemy_2" id: orm.Mapped[int] = orm.mapped_column( BigInteger, Identity(start=0), primary_key=True ) - - Base.metadata.drop_all(engine_obj) + + Base.metadata.drop_all(engine_obj) Base.metadata.create_all(engine_obj) diff --git a/tests/test_orm_dao.py b/tests/test_orm_dao.py index fbcace3..a11fe5d 100644 --- a/tests/test_orm_dao.py +++ b/tests/test_orm_dao.py @@ -1,15 +1,20 @@ -import os -import sys -sys.path.insert(0, 'pysqream_sqlalchemy') -sys.path.insert(0, 'tests') + import pytest import sqlalchemy as sa -from test_base import TestBaseOrm -from sqlalchemy import select, dialects, Table, Column, Identity, ForeignKey, update, delete +from sqlalchemy import ( + Column, + ForeignKey, + Identity, + Table, + delete, + dialects, + select, + update, +) from sqlalchemy.orm import Session +from tests.test_base import TestBaseOrm - -dialects.registry.register("pysqream.dialect", "dialect", "SqreamDialect") +dialects.registry.register("pysqream.dialect", "pysqream_sqlalchemy.dialect", "SqreamDialect") class TestOrmDao(TestBaseOrm): @@ -22,8 +27,8 @@ def test_create_all(self): def test_create_table_with_identity(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, Identity(start=0)), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, Identity(start=0)), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -31,8 +36,8 @@ def test_create_table_with_identity(self): def test_create_table_with_identity_minvalue_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, Identity(start=1, minvalue=1)), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, Identity(start=1, minvalue=1)), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -44,8 +49,8 @@ def test_create_table_with_identity_minvalue_not_supported(self): def test_create_table_with_identity_maxvalue_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, Identity(start=1, maxvalue=10)), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, Identity(start=1, maxvalue=10)), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -57,8 +62,8 @@ def test_create_table_with_identity_maxvalue_not_supported(self): def test_create_table_with_identity_nomaxvalue_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, Identity(start=1, nomaxvalue=10)), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, Identity(start=1, nomaxvalue=10)), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -70,8 +75,8 @@ def test_create_table_with_identity_nomaxvalue_not_supported(self): def test_create_table_with_identity_nominvalue_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, Identity(start=1, nominvalue=10)), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, Identity(start=1, nominvalue=10)), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -83,8 +88,8 @@ def test_create_table_with_identity_nominvalue_not_supported(self): def test_create_table_with_identity_cache_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, Identity(start=1, cache=5)), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, Identity(start=1, cache=5)), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -96,8 +101,8 @@ def test_create_table_with_identity_cache_not_supported(self): def test_create_table_with_identity_order_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, Identity(start=1, order=True)), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, Identity(start=1, order=True)), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -109,8 +114,8 @@ def test_create_table_with_identity_order_not_supported(self): def test_create_table_with_identity_cycle_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, Identity(start=1, cycle=6)), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, Identity(start=1, cycle=6)), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -122,8 +127,8 @@ def test_create_table_with_identity_cycle_not_supported(self): def test_foreign_key_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, ForeignKey("table1.id")), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, ForeignKey("table1.id")), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) @@ -135,8 +140,8 @@ def test_foreign_key_not_supported(self): def test_primary_key_not_supported(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer, primary_key=True), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer, primary_key=True), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) if self.insp.has_table(table3.name): table3.drop(bind=self.engine) diff --git a/tests/test_orm_dto.py b/tests/test_orm_dto.py index 6f50d85..57e2f2c 100644 --- a/tests/test_orm_dto.py +++ b/tests/test_orm_dto.py @@ -1,18 +1,32 @@ -import os -import sys -sys.path.insert(0, 'pysqream_sqlalchemy') -sys.path.insert(0, 'tests') + +from datetime import datetime + import pytest import sqlalchemy as sa -from test_base import TestBaseOrm -from sqlalchemy import select, dialects, Table, Column, union_all, func, distinct, case, cast, Numeric, \ - extract, nulls_first, desc, nulls_last, asc, true, false +from sqlalchemy import ( + Column, + Numeric, + Table, + asc, + case, + cast, + desc, + dialects, + distinct, + extract, + false, + func, + nulls_first, + nulls_last, + select, + true, + union_all, +) +from sqlalchemy.orm import Session, aliased from sqlalchemy.sql import exists -from sqlalchemy.orm import aliased, Session -from datetime import datetime - +from tests.test_base import TestBaseOrm -dialects.registry.register("pysqream.dialect", "dialect", "SqreamDialect") +dialects.registry.register("pysqream.dialect", "pysqream_sqlalchemy.dialect", "SqreamDialect") class TestOrmDto(TestBaseOrm): @@ -36,7 +50,7 @@ def test_select_one_col(self): def test_select_where_not_supported(self): with pytest.raises(Exception) as e_info: - stmt = self.table1.select().where(self.table1.c.id == '1') + stmt = self.table1.select().where(self.table1.c.id == "1") self.session.execute(stmt) assert "Where clause of parameterized query not supported on SQream" in str(e_info.value) @@ -70,7 +84,7 @@ def test_select_group_by(self): def test_grouping_sets(self): stmt = select( - func.sum(self.user.id), self.user.name, self.user.fullname + func.sum(self.user.id), self.user.name, self.user.fullname, ).group_by(func.grouping_sets(self.user.name, self.user.fullname)) with pytest.raises(Exception) as e_info: with Session(self.engine) as session: @@ -91,7 +105,7 @@ def test_select_join_where_not_supported(self): join_stmt = self.table1.join(self.table2, self.table1.c.id == self.table2.c.id) with pytest.raises(Exception) as e_info: - stmt = select(self.table1).select_from(join_stmt).where(self.table1.c.id == '1') + stmt = select(self.table1).select_from(join_stmt).where(self.table1.c.id == "1") self.session.execute(stmt) assert "Where clause of parameterized query not supported on SQream" in str(e_info.value), e_info.value @@ -106,13 +120,13 @@ def test_select_join_order_by(self): def test_select_multiple_join(self): table3 = Table( - 'table3', self.metadata, - Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table3", self.metadata, + Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) table4 = Table( - 'table4', self.metadata, - Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value", sa.Integer) + "table4", self.metadata, + Column("id", sa.Integer), Column("name", sa.UnicodeText), Column("value", sa.Integer), ) expected_stmt = f"SELECT {self.table1.c.id}, {self.table1.c.name}, {self.table1.c.value}, " \ @@ -123,7 +137,7 @@ def test_select_multiple_join(self): f"JOIN {table4.name} ON {self.table1.c.id} = {table4.c.id}" stmt = select(self.table1, self.table2).join_from( - self.table1, self.table2, self.table1.c.id == self.table2.c.id + self.table1, self.table2, self.table1.c.id == self.table2.c.id, ).join_from(self.table1, table3, self.table1.c.id == table3.c.id)\ .join_from(self.table1, table4, self.table1.c.id == table4.c.id) assert expected_stmt == str(stmt) @@ -209,10 +223,10 @@ def test_case_not_supported(self): stmt = select(self.user). \ where( case( - (self.user.name == 'spongebob', 'S'), - (self.user.name == 'jack', 'J'), - else_='E' - ) + (self.user.name == "spongebob", "S"), + (self.user.name == "jack", "J"), + else_="E", + ), ) with pytest.raises(Exception) as e_info: @@ -316,7 +330,7 @@ def test_array_agg(self): def test_char_length_not_supported(self): - stmt = select(func.char_length('daniel')) + stmt = select(func.char_length("daniel")) with pytest.raises(Exception) as e_info: with Session(self.engine) as session: session.execute(stmt) @@ -334,7 +348,7 @@ def test_coalesce_not_supported(self): def test_cube_not_supported(self): - stmt = select(func.sum(self.user.id), self.user.name, self.user.fullname + stmt = select(func.sum(self.user.id), self.user.name, self.user.fullname, ).group_by(func.cube(self.user.name, self.user.fullname)) with pytest.raises(Exception) as e_info: with Session(self.engine) as session: @@ -352,7 +366,7 @@ def test_current_user_not_supported(self): def test_concat_not_supported(self): - stmt = select(func.concat('a', 'b')) + stmt = select(func.concat("a", "b")) with pytest.raises(Exception) as e_info: with Session(self.engine) as session: session.execute(stmt) @@ -397,7 +411,7 @@ def test_cte(self): test_cte = select( self.user.name, - func.sum(self.user.id).label('total_ids') + func.sum(self.user.id).label("total_ids"), ).group_by(self.user.name).cte("test_cte") stmt = select(test_cte) @@ -405,7 +419,7 @@ def test_cte(self): with Session(self.engine) as session: result = session.execute(stmt).all() - assert ('patrick', 3) == result[0], f"expected to get {('patrick', 3)}, got {result[0]}" + assert result[0] == ("patrick", 3), f"expected to get {('patrick', 3)}, got {result[0]}" # sql.expression def test_true(self): From b43dda5e94cd6af8f206bb62b64c7d3756529532 Mon Sep 17 00:00:00 2001 From: Michael Rogozin Date: Tue, 1 Oct 2024 01:45:18 +0300 Subject: [PATCH 2/3] fix: resolve import `Operations` from `alembic` --- tests/test_dialect.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/tests/test_dialect.py b/tests/test_dialect.py index c80062e..c4d48d2 100644 --- a/tests/test_dialect.py +++ b/tests/test_dialect.py @@ -10,11 +10,10 @@ import pytest import sqlalchemy as sa from sqlalchemy import create_engine, select, Table, Column, insert, text, DDL, orm, Identity, BigInteger, MetaData, engine -from test_base import TestBase, Logger, TestBaseTI +from alembic.operations import Operations from alembic.runtime.migration import MigrationContext -from sqlalchemy import DDL, Column, Table, create_engine, insert, orm, select, text +from tests.test_base import TestBase, Logger, TestBaseTI -from tests.test_base import Logger, TestBase, TestBaseTI # Registering dialect sa.dialects.registry.register("pysqream.dialect", "dialect", "SqreamDialect") From 3b85e2e8867b72a26622b60bab2210185b86cb36 Mon Sep 17 00:00:00 2001 From: Michael Rogozin Date: Tue, 1 Oct 2024 15:46:05 +0300 Subject: [PATCH 3/3] remove: less comment --- tests/test_crud.py | 16 ---------------- 1 file changed, 16 deletions(-) diff --git a/tests/test_crud.py b/tests/test_crud.py index 0ee17b6..085bde7 100644 --- a/tests/test_crud.py +++ b/tests/test_crud.py @@ -228,22 +228,6 @@ def test_read_from_table(self, crud_table, statement, expected_result): assert len(result1) == expected_result == len(result2) - # def test_read_from_table_parametrized(self, crud_table): - # self.recreate_all_via_metadata() - # - # self.insert_values_into_table(crud_table, rows_amount=10) - # - # # Write the select query using SQLAlchemy - # query = select([crud_table.c.i]).where(crud_table.c.i == text(":value")) - # - # # Execute the query with parameter binding - # result = connection.execute(query, value=123) # Replace 123 with actual value - # - # # Fetch and print results - # for row in result: - # print(row['i']) - - @pytest.mark.parametrize( ("statement", "expected_result"), (