mirror of
https://github.com/GNS3/gns3-server.git
synced 2026-10-01 00:02:34 +03:00
fix(typing): resolve mypy errors in db.repositories.llm_model_configs
This commit is contained in:
parent
b1d13edd2d
commit
c048124c21
@ -16,12 +16,14 @@
|
||||
# along with this program. If not, see <http://www.gnu.org/licenses/>.
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from fastapi.encoders import jsonable_encoder
|
||||
from sqlalchemy import Column, DateTime, func, inspect
|
||||
from sqlalchemy.types import TypeDecorator, CHAR, VARCHAR
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
@ -105,7 +107,9 @@ class BaseTable(Base):
|
||||
__abstract__ = True
|
||||
|
||||
created_at = Column(DateTime, server_default=func.current_timestamp())
|
||||
updated_at = Column(DateTime, server_default=func.current_timestamp(), onupdate=func.current_timestamp())
|
||||
updated_at: Mapped[Optional[datetime]] = mapped_column(
|
||||
DateTime, server_default=func.current_timestamp(), onupdate=func.current_timestamp()
|
||||
)
|
||||
|
||||
__mapper_args__ = {"eager_defaults": True}
|
||||
|
||||
|
||||
@ -42,7 +42,7 @@ class LLMModelConfig(BaseTable):
|
||||
user_id = Column(GUID, ForeignKey("users.user_id", ondelete="CASCADE"), nullable=True)
|
||||
group_id = Column(GUID, ForeignKey("user_groups.user_group_id", ondelete="CASCADE"), nullable=True)
|
||||
is_default = Column(Boolean, default=False, nullable=False)
|
||||
version = Column(Integer, default=0, nullable=False) # Optimistic locking version
|
||||
version: Mapped[int] = mapped_column(Integer, default=0, nullable=False) # Optimistic locking version
|
||||
|
||||
# Relationships
|
||||
user = relationship("User", backref="llm_model_configs")
|
||||
|
||||
@ -27,6 +27,7 @@ from gns3server.config import Config
|
||||
from gns3server.services import auth_service
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@ -75,7 +76,7 @@ def create_default_super_admin(target, connection, **kw):
|
||||
class UserGroup(BaseTable):
|
||||
__tablename__ = "user_groups"
|
||||
|
||||
user_group_id = Column(GUID, primary_key=True, default=generate_uuid)
|
||||
user_group_id: Mapped[uuid.UUID] = mapped_column(GUID, primary_key=True, default=generate_uuid)
|
||||
name = Column(String, unique=True, index=True)
|
||||
is_builtin = Column(Boolean, default=False)
|
||||
users = relationship("User", secondary=user_group_map, back_populates="groups")
|
||||
|
||||
@ -16,8 +16,9 @@
|
||||
# along with this program. If not, see <http://www.gnu.org/licenses/>.
|
||||
|
||||
from uuid import UUID
|
||||
from typing import Optional, List, Dict, Any
|
||||
from typing import Optional, List, Dict, Any, cast
|
||||
from sqlalchemy import select, update, delete, and_
|
||||
from sqlalchemy.engine import CursorResult
|
||||
from datetime import datetime
|
||||
|
||||
import logging
|
||||
@ -63,7 +64,7 @@ class LLMModelConfigsRepository(BaseRepository):
|
||||
.order_by(models.LLMModelConfig.created_at)
|
||||
)
|
||||
result = await self._db_session.execute(query)
|
||||
return result.scalars().all()
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def get_user_default_config(self, user_id: UUID) -> Optional[models.LLMModelConfig]:
|
||||
"""Get a user's default LLM model configuration."""
|
||||
@ -182,7 +183,7 @@ class LLMModelConfigsRepository(BaseRepository):
|
||||
)
|
||||
result = await self._db_session.execute(query)
|
||||
await self._db_session.commit()
|
||||
return result.rowcount > 0
|
||||
return cast(CursorResult, result).rowcount > 0
|
||||
|
||||
async def set_user_default_config(self, user_id: UUID, config_id: UUID) -> bool:
|
||||
"""Set a user's default LLM model configuration."""
|
||||
@ -202,7 +203,7 @@ class LLMModelConfigsRepository(BaseRepository):
|
||||
)
|
||||
result = await self._db_session.execute(query)
|
||||
await self._db_session.commit()
|
||||
return result.rowcount > 0
|
||||
return cast(CursorResult, result).rowcount > 0
|
||||
|
||||
# Group configuration methods
|
||||
|
||||
@ -222,7 +223,7 @@ class LLMModelConfigsRepository(BaseRepository):
|
||||
.order_by(models.LLMModelConfig.created_at)
|
||||
)
|
||||
result = await self._db_session.execute(query)
|
||||
return result.scalars().all()
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def get_group_default_config(self, group_id: UUID) -> Optional[models.LLMModelConfig]:
|
||||
"""Get a group's default LLM model configuration."""
|
||||
@ -341,7 +342,7 @@ class LLMModelConfigsRepository(BaseRepository):
|
||||
)
|
||||
result = await self._db_session.execute(query)
|
||||
await self._db_session.commit()
|
||||
return result.rowcount > 0
|
||||
return cast(CursorResult, result).rowcount > 0
|
||||
|
||||
async def set_group_default_config(self, group_id: UUID, config_id: UUID) -> bool:
|
||||
"""Set a group's default LLM model configuration."""
|
||||
@ -361,7 +362,7 @@ class LLMModelConfigsRepository(BaseRepository):
|
||||
)
|
||||
result = await self._db_session.execute(query)
|
||||
await self._db_session.commit()
|
||||
return result.rowcount > 0
|
||||
return cast(CursorResult, result).rowcount > 0
|
||||
|
||||
# Inheritance methods
|
||||
|
||||
@ -441,21 +442,21 @@ class LLMModelConfigsRepository(BaseRepository):
|
||||
)
|
||||
|
||||
# Select default_config with proper priority:
|
||||
# 1. User's config marked with is_default: true
|
||||
# 2. Group's config marked with is_default: true
|
||||
# 3. First config in the list (user configs come first)
|
||||
for config in configs_with_source:
|
||||
if config["is_default"] and config["source"] == "user":
|
||||
default_config = config
|
||||
# 1. User's entry marked with is_default: true
|
||||
# 2. Group's entry marked with is_default: true
|
||||
# 3. First entry in the list (user configs come first)
|
||||
for entry in configs_with_source:
|
||||
if entry["is_default"] and entry["source"] == "user":
|
||||
default_config = entry
|
||||
break
|
||||
|
||||
if default_config is None:
|
||||
for config in configs_with_source:
|
||||
if config["is_default"] and config["source"] == "group":
|
||||
default_config = config
|
||||
for entry in configs_with_source:
|
||||
if entry["is_default"] and entry["source"] == "group":
|
||||
default_config = entry
|
||||
break
|
||||
|
||||
# Fallback to first config if no default is marked
|
||||
# Fallback to first entry if no default is marked
|
||||
if default_config is None and configs_with_source:
|
||||
default_config = configs_with_source[0]
|
||||
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user