Fix typing error: session type

This commit is contained in:
John Davis
2024-04-02 10:08:53 -04:00
parent 79261dd1f4
commit 78aa9709b2
11 changed files with 28 additions and 37 deletions
+3 -4
View File
@@ -11,7 +11,6 @@ from sqlalchemy import (
)
from sqlalchemy.dialects.postgresql import insert as ps_insert
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from galaxy.model import CeleryUserRateLimit
from galaxy.model.base import transaction
@@ -70,7 +69,7 @@ class GalaxyTaskBeforeStartUserRateLimit(GalaxyTaskBeforeStart):
@abstractmethod
def calculate_task_start_time(
self, user_id: int, sa_session: Session, task_interval_secs: float, now: datetime.datetime
self, user_id: int, sa_session: galaxy_scoped_session, task_interval_secs: float, now: datetime.datetime
) -> datetime.datetime:
return now
@@ -99,7 +98,7 @@ class GalaxyTaskBeforeStartUserRateLimitPostgres(GalaxyTaskBeforeStartUserRateLi
)
def calculate_task_start_time( # type: ignore
self, user_id: int, sa_session: Session, task_interval_secs: float, now: datetime.datetime
self, user_id: int, sa_session: galaxy_scoped_session, task_interval_secs: float, now: datetime.datetime
) -> datetime.datetime:
with transaction(sa_session):
result = sa_session.execute(
@@ -138,7 +137,7 @@ class GalaxyTaskBeforeStartUserRateLimitStandard(GalaxyTaskBeforeStartUserRateLi
)
def calculate_task_start_time(
self, user_id: int, sa_session: Session, task_interval_secs: float, now: datetime.datetime
self, user_id: int, sa_session: galaxy_scoped_session, task_interval_secs: float, now: datetime.datetime
) -> datetime.datetime:
last_scheduled_time = None
with transaction(sa_session):
+2 -2
View File
@@ -14,9 +14,9 @@ from typing import (
)
from sqlalchemy import select
from sqlalchemy.orm import Session
from galaxy.model import HistoryDatasetAssociation
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.util import (
galaxy_directory,
sanitize_lists_to_string,
@@ -166,6 +166,6 @@ class GenomeBuilds:
return (chrom_info, db_dataset)
def get_len_files_by_history(session: Session, history_id: int):
def get_len_files_by_history(session: galaxy_scoped_session, history_id: int):
stmt = select(HistoryDatasetAssociation).filter_by(history_id=history_id, extension="len", deleted=False)
return session.scalars(stmt)
+2 -2
View File
@@ -5,13 +5,13 @@ from typing import (
)
from sqlalchemy import select
from sqlalchemy.orm import Session
from galaxy import model
from galaxy.exceptions import ObjectNotFound
from galaxy.managers.context import ProvidesAppContext
from galaxy.model import GroupRoleAssociation
from galaxy.model.base import transaction
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.structured_app import MinimalManagerApp
log = logging.getLogger(__name__)
@@ -93,7 +93,7 @@ class GroupRolesManager:
trans.sa_session.commit()
def get_group_role(session: Session, group, role) -> Optional[GroupRoleAssociation]:
def get_group_role(session: galaxy_scoped_session, group, role) -> Optional[GroupRoleAssociation]:
stmt = (
select(GroupRoleAssociation).where(GroupRoleAssociation.group == group).where(GroupRoleAssociation.role == role)
)
+2 -2
View File
@@ -5,7 +5,6 @@ from typing import (
)
from sqlalchemy import select
from sqlalchemy.orm import Session
from galaxy import model
from galaxy.exceptions import ObjectNotFound
@@ -15,6 +14,7 @@ from galaxy.model import (
UserGroupAssociation,
)
from galaxy.model.base import transaction
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.structured_app import MinimalManagerApp
log = logging.getLogger(__name__)
@@ -96,7 +96,7 @@ class GroupUsersManager:
trans.sa_session.commit()
def get_group_user(session: Session, user, group) -> Optional[UserGroupAssociation]:
def get_group_user(session: galaxy_scoped_session, user, group) -> Optional[UserGroupAssociation]:
stmt = (
select(UserGroupAssociation).where(UserGroupAssociation.user == user).where(UserGroupAssociation.group == group)
)
+2 -3
View File
@@ -2,7 +2,6 @@ from sqlalchemy import (
false,
select,
)
from sqlalchemy.orm import Session
from galaxy import model
from galaxy.exceptions import (
@@ -152,11 +151,11 @@ class GroupsManager:
return group
def get_group_by_name(session: Session, name: str):
def get_group_by_name(session: galaxy_scoped_session, name: str):
stmt = select(Group).filter(Group.name == name).limit(1)
return session.scalars(stmt).first()
def get_not_deleted_groups(session: Session):
def get_not_deleted_groups(session: galaxy_scoped_session):
stmt = select(Group).where(Group.deleted == false())
return session.scalars(stmt)
+2 -5
View File
@@ -20,10 +20,7 @@ from sqlalchemy import (
or_,
true,
)
from sqlalchemy.orm import (
aliased,
Session,
)
from sqlalchemy.orm import aliased
from sqlalchemy.sql import select
from galaxy import model
@@ -1069,7 +1066,7 @@ def summarize_job_outputs(job: model.Job, tool, params):
return outputs
def get_jobs_to_check_at_startup(session: Session, track_jobs_in_database: bool, config):
def get_jobs_to_check_at_startup(session: galaxy_scoped_session, track_jobs_in_database: bool, config):
if track_jobs_in_database:
in_list = (Job.states.QUEUED, Job.states.RUNNING, Job.states.STOPPED)
else:
+6 -8
View File
@@ -24,10 +24,7 @@ from sqlalchemy import (
select,
true,
)
from sqlalchemy.orm import (
aliased,
Session,
)
from sqlalchemy.orm import aliased
from galaxy import (
exceptions,
@@ -64,6 +61,7 @@ from galaxy.model.index_filter_util import (
text_column_filter,
)
from galaxy.model.item_attrs import UsesAnnotations
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.schema.schema import (
CreatePagePayload,
PageContentFormat,
@@ -644,12 +642,12 @@ def placeholderRenderForSave(trans: ProvidesHistoryContext, item_class, item_id,
)
def get_page_revision(session: Session, page_id: int):
def get_page_revision(session: galaxy_scoped_session, page_id: int):
stmt = select(PageRevision).filter_by(page_id=page_id)
return session.scalars(stmt)
def get_shared_pages(session: Session, user: User):
def get_shared_pages(session: galaxy_scoped_session, user: User):
stmt = (
select(PageUserShareAssociation)
.where(PageUserShareAssociation.user == user)
@@ -660,12 +658,12 @@ def get_shared_pages(session: Session, user: User):
return session.scalars(stmt)
def get_page(session: Session, user: User, slug: str):
def get_page(session: galaxy_scoped_session, user: User, slug: str):
stmt = _build_page_query(select(Page), user, slug)
return session.scalars(stmt).first()
def page_exists(session: Session, user: User, slug: str) -> bool:
def page_exists(session: galaxy_scoped_session, user: User, slug: str) -> bool:
stmt = _build_page_query(select(Page.id), user, slug)
return session.scalars(stmt).first() is not None
+3 -5
View File
@@ -9,10 +9,7 @@ from sqlalchemy import (
false,
select,
)
from sqlalchemy.orm import (
exc as sqlalchemy_exceptions,
Session,
)
from sqlalchemy.orm import exc as sqlalchemy_exceptions
from galaxy import model
from galaxy.exceptions import (
@@ -26,6 +23,7 @@ from galaxy.managers import base
from galaxy.managers.context import ProvidesUserContext
from galaxy.model import Role
from galaxy.model.base import transaction
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.schema.schema import RoleDefinitionModel
from galaxy.util import unicodify
@@ -162,6 +160,6 @@ class RoleManager(base.ModelManager[model.Role]):
return role
def get_roles_by_ids(session: Session, role_ids):
def get_roles_by_ids(session: galaxy_scoped_session, role_ids):
stmt = select(Role).where(Role.id.in_(role_ids))
return session.scalars(stmt).all()
+2 -2
View File
@@ -24,7 +24,6 @@ from sqlalchemy import (
select,
true,
)
from sqlalchemy.orm import Session
from sqlalchemy.orm.exc import NoResultFound
from galaxy import (
@@ -46,6 +45,7 @@ from galaxy.model import (
UserQuotaUsage,
)
from galaxy.model.base import transaction
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.security.validate_user_input import (
VALID_EMAIL_RE,
validate_email,
@@ -873,7 +873,7 @@ class AdminUserFilterParser(base.ModelFilterParser, deletable.PurgableFiltersMix
self.fn_filter_parsers.update({})
def get_users_by_ids(session: Session, user_ids):
def get_users_by_ids(session: galaxy_scoped_session, user_ids):
stmt = select(User).where(User.id.in_(user_ids))
return session.scalars(stmt).all()
@@ -20,7 +20,6 @@ from sqlalchemy import (
select,
true,
)
from sqlalchemy.orm import Session
from galaxy import (
exceptions as glx_exceptions,
@@ -45,6 +44,7 @@ from galaxy.managers.notification import NotificationManager
from galaxy.managers.users import UserManager
from galaxy.model import HistoryDatasetAssociation
from galaxy.model.base import transaction
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.model.store import payload_to_source_uri
from galaxy.schema import (
FilterQueryParams,
@@ -820,7 +820,7 @@ class HistoriesService(ServiceBase, ConsumesModelStores, ServesExportStores):
return None
def get_fasta_hdas_by_history(session: Session, history_id: int):
def get_fasta_hdas_by_history(session: galaxy_scoped_session, history_id: int):
stmt = (
select(HistoryDatasetAssociation)
.filter_by(history_id=history_id, extension="fasta", deleted=False)
+2 -2
View File
@@ -6,7 +6,6 @@ from sqlalchemy import (
select,
true,
)
from sqlalchemy.orm import Session
from galaxy import util
from galaxy.managers.context import ProvidesUserContext
@@ -14,6 +13,7 @@ from galaxy.managers.groups import get_group_by_name
from galaxy.managers.quotas import QuotaManager
from galaxy.managers.users import get_user_by_email
from galaxy.model import Quota
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.quota._schema import (
CreateQuotaParams,
CreateQuotaResult,
@@ -161,7 +161,7 @@ class QuotasService(ServiceBase):
payload["in_groups"] = list(map(str, new_in_groups))
def get_quotas(session: Session, deleted: bool = False):
def get_quotas(session: galaxy_scoped_session, deleted: bool = False):
is_deleted = true()
if not deleted:
is_deleted = false()