From 6bfdb4eb841fab8f9cd062eaac57b2dd0cdf6af4 Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Sat, 6 Jun 2026 23:30:16 +0800 Subject: [PATCH 01/10] Added database url to env.example --- .env.example | 3 ++- Database.db | Bin 0 -> 20480 bytes 2 files changed, 2 insertions(+), 1 deletion(-) create mode 100644 Database.db diff --git a/.env.example b/.env.example index 20a9ad8..d22255c 100644 --- a/.env.example +++ b/.env.example @@ -1,3 +1,4 @@ GEMINI_API_KEY=your_api_key_here GEMINI_MODEL=gemini-2.5-flash -MCP_SERVER_URL=http://localhost:8003/mcp \ No newline at end of file +MCP_SERVER_URL=http://localhost:8003/mcp +DATABASE_URL = "sqlite:///./Database.db" \ No newline at end of file diff --git a/Database.db b/Database.db new file mode 100644 index 0000000000000000000000000000000000000000..31cca35887f3b373887e970066a78a24027d8bb0 GIT binary patch literal 20480 zcmeI%zi!h&9Ki8&nzl4j`e(|Rld>cdLJ254ZaEbT*P+Hi>SVdGkqk|o#)kymD9^xy z@Cua}cq2x90T*WoVquBCXZifQ`&pKMpS#8Kt8+KiiF}hSrjeFU#G$Y(@l;A7M6Ect zi?eKITrUU3)vCnn!kT#W<$Lqshp0F2#QxWVZ_T^?n{_ucj{pJ)Ab)l#RVvG)?4d+dp}3``yRGVQ(d?KhFPT2mStP_NDKf+5Uxmp)O=MjeF*KW6wDs zE7>*A?KOJC`cBK~SR$Rp$%p*F0;hJ(sr`}WA%ZMECh=j&hPM)|ui zj&yPS4r;5Vmec>K&^XB_i*&BjY<6!|X!MsAPqOLUtVk@6MrB|!mMhFissdSG@{({zWXqalC*0tg_000IagfB*sr zAb`N83N%b%eE)Ci@=_N92q1s}0tg_000IagfB*tZ0sjB95fDHC0R#|0009ILKmY** g5ZHVH{{L_O8dD(x2q1s}0tg_000IagfB*tN0XFQi=>Px# literal 0 HcmV?d00001 From 9bee48b297bd2c760a6f49d503e500f47805a888 Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Sat, 6 Jun 2026 23:53:53 +0800 Subject: [PATCH 02/10] Add type hints and update docstrings for CRUD methods --- app/services/database.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/app/services/database.py b/app/services/database.py index 314e0f1..2413496 100644 --- a/app/services/database.py +++ b/app/services/database.py @@ -25,19 +25,19 @@ class DataBaseMethods: # Adds a data object to database @staticmethod - def add_object(db: Session, new_object): + def add_object(db: Session, new_object: DBTask) -> DBTask: """ Adds a database object to the database. Args: db (Session): Database session, can be grabbed with get_db() - new_object (database_objects): New + new_object (DBTask): New database object to be added must be inside the database_objects.py module Returns: - database_object: returns the entered object + DBTask: returns the entered object if it was successfully added """ try: @@ -54,13 +54,13 @@ def add_object(db: Session, new_object): # deletes a data object from database @staticmethod - def delete_object(db: Session, object_to_delete): + def delete_object(db: Session, object_to_delete: DBTask) -> bool: """ Deletes a database object from the database. Args: db (Session): Database session, can be grabbed with get_db() - object_to_delete (database_objects): New database object to be added + object_to_delete (DBTask): New database object to be added must be inside the database_objects.py module Returns: @@ -80,18 +80,18 @@ def delete_object(db: Session, object_to_delete): # Gets an object by Name from the database @staticmethod - def get_object_by_name(db: Session, obj_ref, name: str): + def get_object_by_name(db: Session, obj_ref, name: str) -> DBTask | None: """ Gets an object based on the name entered. Args: db (Session): Database session, can be grabbed with get_db() name (string): Name of the object - obj_ref (database_objects): Object reference, so an object from + obj_ref (DBTask): Object reference, so an object from the database_objects.py module Returns: - database_object: returns the object if it was successfully found + DBTask: returns the object if it was successfully found, otherwise None """ if not obj_ref.name: raise HTTPException(status_code=500, @@ -106,18 +106,18 @@ def get_object_by_name(db: Session, obj_ref, name: str): # Gets an object by ID from the database, data object must have ID field @staticmethod - def get_object_by_id(db: Session, obj_ref, object_id: int): + def get_object_by_id(db: Session, obj_ref, object_id: int)-> DBTask | None: """ Gets an object based on the id entered. Args: db (Session): Database session, can be grabbed with get_db() object_id (int): ID of the object - obj_ref (database_objects): Object reference, so an object from + obj_ref (DBTask): Object reference, so an object from the database_objects.py module Returns: - database_object: returns the object if it was successfully found + DBTask: returns the object if it was successfully found """ try: param = obj_ref.id @@ -127,7 +127,7 @@ def get_object_by_id(db: Session, obj_ref, object_id: int): detail=f"Object with id: {id}, does not exist") @staticmethod - def query_db(db: Session, obj_ref, field: str, value): + def query_db(db: Session, obj_ref, field: str, value) -> list[DBTask]: if value is None: try: return db.query(obj_ref).filter(obj_ref.id).all() From 06e856e19ab2034ca12292d73d1e03d725d63899 Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Fri, 19 Jun 2026 10:06:36 +0800 Subject: [PATCH 03/10] fix: remove Database.db from version control and add to .gitignore --- .gitignore | 1 + Database.db | Bin 20480 -> 0 bytes 2 files changed, 1 insertion(+) delete mode 100644 Database.db diff --git a/.gitignore b/.gitignore index 4e5c117..55c131e 100644 --- a/.gitignore +++ b/.gitignore @@ -161,3 +161,4 @@ cython_debug/ .DS_Store .ruff_cache .zed +Database.db diff --git a/Database.db b/Database.db deleted file mode 100644 index 31cca35887f3b373887e970066a78a24027d8bb0..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 20480 zcmeI%zi!h&9Ki8&nzl4j`e(|Rld>cdLJ254ZaEbT*P+Hi>SVdGkqk|o#)kymD9^xy z@Cua}cq2x90T*WoVquBCXZifQ`&pKMpS#8Kt8+KiiF}hSrjeFU#G$Y(@l;A7M6Ect zi?eKITrUU3)vCnn!kT#W<$Lqshp0F2#QxWVZ_T^?n{_ucj{pJ)Ab)l#RVvG)?4d+dp}3``yRGVQ(d?KhFPT2mStP_NDKf+5Uxmp)O=MjeF*KW6wDs zE7>*A?KOJC`cBK~SR$Rp$%p*F0;hJ(sr`}WA%ZMECh=j&hPM)|ui zj&yPS4r;5Vmec>K&^XB_i*&BjY<6!|X!MsAPqOLUtVk@6MrB|!mMhFissdSG@{({zWXqalC*0tg_000IagfB*sr zAb`N83N%b%eE)Ci@=_N92q1s}0tg_000IagfB*tZ0sjB95fDHC0R#|0009ILKmY** g5ZHVH{{L_O8dD(x2q1s}0tg_000IagfB*tN0XFQi=>Px# From ffa7ff76f0068e391c09a45ca46204d6f976faca Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Fri, 19 Jun 2026 10:09:07 +0800 Subject: [PATCH 04/10] fix: add missing imports and load DATABASE_URL from env --- app/services/database.py | 50 +++++++++++++++++++++++++--------------- 1 file changed, 31 insertions(+), 19 deletions(-) diff --git a/app/services/database.py b/app/services/database.py index 2413496..3455fd9 100644 --- a/app/services/database.py +++ b/app/services/database.py @@ -1,8 +1,16 @@ +from __future__ import annotations + +import os +from typing import TYPE_CHECKING + from fastapi import HTTPException from sqlalchemy import create_engine from sqlalchemy.orm import Session, declarative_base, sessionmaker -Database_url = "sqlite:///./Database.db" +if TYPE_CHECKING: + from app.data.database_objects import DBTask + +Database_url = os.getenv("DATABASE_URL", "sqlite:///./Database.db") Engine = create_engine(Database_url, connect_args={"check_same_thread": False}) LocalSession = sessionmaker(bind=Engine, autocommit=False, autoflush=False) Base = declarative_base() @@ -49,8 +57,9 @@ def add_object(db: Session, new_object: DBTask) -> DBTask: return new_object except Exception as err: db.rollback() - raise HTTPException(status_code=404, - detail="Object already exists") from err + raise HTTPException( + status_code=404, detail="Object already exists" + ) from err # deletes a data object from database @staticmethod @@ -75,8 +84,9 @@ def delete_object(db: Session, object_to_delete: DBTask) -> bool: return True except Exception as err: db.rollback() - raise HTTPException(status_code=404, - detail=f"Object does not exist, {err}") from err + raise HTTPException( + status_code=404, detail=f"Object does not exist, {err}" + ) from err # Gets an object by Name from the database @staticmethod @@ -94,19 +104,19 @@ def get_object_by_name(db: Session, obj_ref, name: str) -> DBTask | None: DBTask: returns the object if it was successfully found, otherwise None """ if not obj_ref.name: - raise HTTPException(status_code=500, - detail="Object must have a name") + raise HTTPException(status_code=500, detail="Object must have a name") try: param = obj_ref.name return db.query(obj_ref).filter(param == name).first() except Exception: - HTTPException(status_code=404, - detail=f"Object with name: {name}, does not exist") + HTTPException( + status_code=404, detail=f"Object with name: {name}, does not exist" + ) # Gets an object by ID from the database, data object must have ID field @staticmethod - def get_object_by_id(db: Session, obj_ref, object_id: int)-> DBTask | None: + def get_object_by_id(db: Session, obj_ref, object_id: int) -> DBTask | None: """ Gets an object based on the id entered. @@ -123,8 +133,9 @@ def get_object_by_id(db: Session, obj_ref, object_id: int)-> DBTask | None: param = obj_ref.id return db.query(obj_ref).filter(param == object_id).first() except Exception: - HTTPException(status_code=404, - detail=f"Object with id: {id}, does not exist") + HTTPException( + status_code=404, detail=f"Object with id: {id}, does not exist" + ) @staticmethod def query_db(db: Session, obj_ref, field: str, value) -> list[DBTask]: @@ -132,17 +143,19 @@ def query_db(db: Session, obj_ref, field: str, value) -> list[DBTask]: try: return db.query(obj_ref).filter(obj_ref.id).all() except Exception: - HTTPException(status_code=404, - detail=f"Object with filed: {field}," - f" does not exist") + HTTPException( + status_code=404, + detail=f"Object with filed: {field}, does not exist", + ) else: try: param = getattr(obj_ref, field) return db.query(obj_ref).filter(param == value).all() except Exception: - HTTPException(status_code=404, - detail=f"Object with filed: {field}," - f" does not exist") + HTTPException( + status_code=404, + detail=f"Object with filed: {field}, does not exist", + ) def get_session() -> Session: @@ -167,7 +180,6 @@ def clear_table(session, db_object): print(f"Deleted {deleted} rows") except Exception as e: - session.rollback() print("CLEAR ERROR:", e) From 58a988b94e7c472878651e55ca65e2a0d8229828 Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Fri, 19 Jun 2026 10:09:34 +0800 Subject: [PATCH 05/10] fix: ruff lint and format fixes --- app/api/task_handler.py | 32 ++++++++++---------- app/routes/route_tasks.py | 30 ++++++++++++------- app/routes/router_handler.py | 10 ++----- tests/misc/test_database.py | 52 ++++++++++++++++++++++++-------- tests/misc/test_tasks.py | 57 +++++++++++++++++++++++++++++------- 5 files changed, 124 insertions(+), 57 deletions(-) diff --git a/app/api/task_handler.py b/app/api/task_handler.py index dbf053f..2a4450b 100644 --- a/app/api/task_handler.py +++ b/app/api/task_handler.py @@ -22,24 +22,26 @@ class TaskHandler: """ @classmethod - def create_task(cls, name: str, _type: str, - description: str, completed: bool = False) -> PLTask: - return PLTask(name=name, - type=_type, - description=description, - completed=completed) + def create_task( + cls, name: str, _type: str, description: str, completed: bool = False + ) -> PLTask: + return PLTask( + name=name, type=_type, description=description, completed=completed + ) # creates a task with curl content @classmethod def add_task_to_db(cls, db: Session, task: PLTask): """Creates a task and adds it to the misc.""" print(db) - new_task_obj = DBTask(id=task.id or None, - name=task.name, - type=task.type, - description=task.description, - completed=task.completed or False, - task_started=task.task_started or datetime.now()) + new_task_obj = DBTask( + id=task.id or None, + name=task.name, + type=task.type, + description=task.description, + completed=task.completed or False, + task_started=task.task_started or datetime.now(), + ) new_task_obj = DataBaseMethods.add_object(db, new_task_obj) @@ -115,8 +117,7 @@ def complete_task(cls, db: Session, task: PLTask): db.commit() return task except Exception as err: - raise HTTPException(status_code=500, - detail=str(err)) from err + raise HTTPException(status_code=500, detail=str(err)) from err else: raise HTTPException(status_code=404, detail="Task not found") @@ -125,8 +126,7 @@ def task_run_duration(cls, task: PLTask): if task.task_started and task.task_ended: time_elapsed = task.task_ended - task.task_started return time_elapsed.total_seconds() - raise HTTPException(status_code=500, - detail="Task hasn't finished running") + raise HTTPException(status_code=500, detail="Task hasn't finished running") @classmethod def __from_db_to_pl(cls, dbtask: DBTask): diff --git a/app/routes/route_tasks.py b/app/routes/route_tasks.py index 9afca1c..fcad947 100644 --- a/app/routes/route_tasks.py +++ b/app/routes/route_tasks.py @@ -12,7 +12,9 @@ @task_router.get("/get-name", response_model=PLTask) -def get_task_by_name(session: Annotated[Session, Depends(get_session_api)], name: str | None = None) -> PLTask: +def get_task_by_name( + session: Annotated[Session, Depends(get_session_api)], name: str | None = None +) -> PLTask: try: task = TaskHandler.get_task(session, name) except Exception as e: @@ -20,15 +22,16 @@ def get_task_by_name(session: Annotated[Session, Depends(get_session_api)], name if task is None: raise HTTPException( - status_code=404, - detail=f"Task with name '{name}' could not be found." + status_code=404, detail=f"Task with name '{name}' could not be found." ) return task @task_router.get("/get-id", response_model=PLTask) -def get_task_by_id(session: Annotated[Session, Depends(get_session_api)], task_id: int | None = None) -> PLTask: +def get_task_by_id( + session: Annotated[Session, Depends(get_session_api)], task_id: int | None = None +) -> PLTask: try: task = TaskHandler.get_task(session, task_id) except Exception as e: @@ -36,20 +39,23 @@ def get_task_by_id(session: Annotated[Session, Depends(get_session_api)], task_i if task is None: raise HTTPException( - status_code=404, - detail=f"Task with id '{task_id}' could not be found." + status_code=404, detail=f"Task with id '{task_id}' could not be found." ) return task @task_router.post("/delete") -def delete_task_by_name(session: Annotated[Session, Depends(get_session_api)], name: str | None = None): +def delete_task_by_name( + session: Annotated[Session, Depends(get_session_api)], name: str | None = None +): TaskHandler.delete_task(session, name) @task_router.post("/add") -def create_task(session: Annotated[Session, Depends(get_session_api)], task: PLTask = None) -> PLTask: +def create_task( + session: Annotated[Session, Depends(get_session_api)], task: PLTask = None +) -> PLTask: try: task = TaskHandler.add_task_to_db(session, task) except Exception as e: @@ -58,11 +64,15 @@ def create_task(session: Annotated[Session, Depends(get_session_api)], task: PLT @task_router.get("/task-list") -def filter_task_by_type(session: Annotated[Session, Depends(get_session_api)], task_type: str | None = None) -> list[PLTask]: +def filter_task_by_type( + session: Annotated[Session, Depends(get_session_api)], task_type: str | None = None +) -> list[PLTask]: return TaskHandler.list_tasks(session, task_type) @task_router.post("/complete") -def complete_task(session: Annotated[Session, Depends(get_session_api)], name: str | None = None) -> PLTask: +def complete_task( + session: Annotated[Session, Depends(get_session_api)], name: str | None = None +) -> PLTask: task = TaskHandler.get_task(session, name) return TaskHandler.complete_task(session, task) diff --git a/app/routes/router_handler.py b/app/routes/router_handler.py index 59c80c4..7024f8c 100644 --- a/app/routes/router_handler.py +++ b/app/routes/router_handler.py @@ -7,7 +7,6 @@ class Router: - # Routers must have the word "_router" in it system_router = APIRouter(prefix="", tags=["system"]) @@ -17,7 +16,6 @@ class Router: def load_all_routes(cls): for root, _dirs, files in os.walk("app/routes"): for file in files: - file_path = os.path.join(root, file) try: if "__pycache__" in file_path: @@ -29,7 +27,7 @@ def load_all_routes(cls): module = module_from_spec(spec) spec.loader.exec_module(module) - except (UnicodeDecodeError, PermissionError): + except UnicodeDecodeError, PermissionError: continue @classmethod @@ -49,11 +47,9 @@ def get_router(cls, router_name): try: return getattr(cls, router_name) except ModuleNotFoundError as err: - raise HTTPException(status_code=404, - detail="Router not found") from err + raise HTTPException(status_code=404, detail="Router not found") from err else: try: return getattr(cls, router_name + "_router") except ModuleNotFoundError as err: - raise HTTPException(status_code=404, - detail="Router not found") from err + raise HTTPException(status_code=404, detail="Router not found") from err diff --git a/tests/misc/test_database.py b/tests/misc/test_database.py index 35d8afe..5b9d465 100644 --- a/tests/misc/test_database.py +++ b/tests/misc/test_database.py @@ -5,7 +5,6 @@ class MyTestCase(unittest.TestCase): - def setUp(self): self.session = LocalSession() @@ -19,33 +18,60 @@ def test_add_to_db(self): def test_get_object_by_name(self): clear_table(self.session, DBTestObject) - db_obj = DataBaseMethods.add_object(self.session, DBTestObject(id=1, name="Test1", type="Flying")) + db_obj = DataBaseMethods.add_object( + self.session, DBTestObject(id=1, name="Test1", type="Flying") + ) - self.assertEqual(db_obj, DataBaseMethods.get_object_by_name(self.session, DBTestObject, "Test1")) + self.assertEqual( + db_obj, + DataBaseMethods.get_object_by_name(self.session, DBTestObject, "Test1"), + ) def test_get_object_by_id(self): clear_table(self.session, DBTestObject) - db_obj = DataBaseMethods.add_object(self.session, DBTestObject(id=1, name="Test1", type="Flying")) + db_obj = DataBaseMethods.add_object( + self.session, DBTestObject(id=1, name="Test1", type="Flying") + ) - self.assertEqual(db_obj, DataBaseMethods.get_object_by_id(self.session, DBTestObject, 1)) + self.assertEqual( + db_obj, DataBaseMethods.get_object_by_id(self.session, DBTestObject, 1) + ) def test_db_query(self): clear_table(self.session, DBTestObject) - db_obj1 = DataBaseMethods.add_object(self.session, DBTestObject(id=1, name="Test1", type="Flying")) - db_obj2 = DataBaseMethods.add_object(self.session, DBTestObject(id=2, name="Test2", type="Ground")) - - self.assertNotEqual([db_obj1, db_obj2], DataBaseMethods.query_db(self.session, DBTestObject, "type", "Flying")) - self.assertEqual([db_obj1], DataBaseMethods.query_db(self.session, DBTestObject, "type", "Flying")) + db_obj1 = DataBaseMethods.add_object( + self.session, DBTestObject(id=1, name="Test1", type="Flying") + ) + db_obj2 = DataBaseMethods.add_object( + self.session, DBTestObject(id=2, name="Test2", type="Ground") + ) + + self.assertNotEqual( + [db_obj1, db_obj2], + DataBaseMethods.query_db(self.session, DBTestObject, "type", "Flying"), + ) + self.assertEqual( + [db_obj1], + DataBaseMethods.query_db(self.session, DBTestObject, "type", "Flying"), + ) def test_object_delete(self): clear_table(self.session, DBTestObject) - db_obj1 = DataBaseMethods.add_object(self.session, DBTestObject(id=1, name="Test1", type="Flying")) - self.assertEqual(db_obj1, DataBaseMethods.get_object_by_name(self.session, DBTestObject, "Test1")) + db_obj1 = DataBaseMethods.add_object( + self.session, DBTestObject(id=1, name="Test1", type="Flying") + ) + self.assertEqual( + db_obj1, + DataBaseMethods.get_object_by_name(self.session, DBTestObject, "Test1"), + ) DataBaseMethods.delete_object(self.session, db_obj1) - self.assertNotEqual(db_obj1, DataBaseMethods.get_object_by_name(self.session, DBTestObject, "Test1")) + self.assertNotEqual( + db_obj1, + DataBaseMethods.get_object_by_name(self.session, DBTestObject, "Test1"), + ) if __name__ == "__main__": diff --git a/tests/misc/test_tasks.py b/tests/misc/test_tasks.py index 0dcf0fa..7ce95db 100644 --- a/tests/misc/test_tasks.py +++ b/tests/misc/test_tasks.py @@ -7,37 +7,72 @@ class TaskTests(unittest.TestCase): - def setUp(self): self.session = LocalSession() self.example_pltasks = [ - PLTask(name="Task1", type="Agriculture", description="Give names for 100 plants", completed=False), - PLTask(name="Task2", type="Construction", description="Give price for new airport", completed=True) + PLTask( + name="Task1", + type="Agriculture", + description="Give names for 100 plants", + completed=False, + ), + PLTask( + name="Task2", + type="Construction", + description="Give price for new airport", + completed=True, + ), ] self.example_dbtask = [ - PLTask(id=1, name="Task1", type="Agriculture", description="Give names for 100 plants", completed=False), - PLTask(id=2, name="Task2", type="Construction", description="Give price for new airport", completed=True) + PLTask( + id=1, + name="Task1", + type="Agriculture", + description="Give names for 100 plants", + completed=False, + ), + PLTask( + id=2, + name="Task2", + type="Construction", + description="Give price for new airport", + completed=True, + ), ] def test_creating_tasks(self): - task1 = TaskHandler.create_task("Task1", "Agriculture", "Give names for 100 plants", False) - task2 = TaskHandler.create_task("Task2", "Construction", "Give price for new airport", True) + task1 = TaskHandler.create_task( + "Task1", "Agriculture", "Give names for 100 plants", False + ) + task2 = TaskHandler.create_task( + "Task2", "Construction", "Give price for new airport", True + ) self.assertEqual(self.example_pltasks[0], task1) self.assertEqual(self.example_pltasks[1], task2) def test_adding_task_to_db(self): clear_table(self.session, DBTask) - self.assertEqual(self.example_dbtask[0], TaskHandler.add_task_to_db(self.session, self.example_pltasks[0])) - self.assertEqual(self.example_dbtask[1], TaskHandler.add_task_to_db(self.session, self.example_pltasks[1])) + self.assertEqual( + self.example_dbtask[0], + TaskHandler.add_task_to_db(self.session, self.example_pltasks[0]), + ) + self.assertEqual( + self.example_dbtask[1], + TaskHandler.add_task_to_db(self.session, self.example_pltasks[1]), + ) def test_finding_task_by_name(self): clear_table(self.session, DBTask) TaskHandler.add_task_to_db(self.session, self.example_pltasks[0]) TaskHandler.add_task_to_db(self.session, self.example_pltasks[1]) - self.assertEqual(self.example_dbtask[0], TaskHandler.get_task(self.session, "Task1")) - self.assertEqual(self.example_dbtask[1], TaskHandler.get_task(self.session, "Task2")) + self.assertEqual( + self.example_dbtask[0], TaskHandler.get_task(self.session, "Task1") + ) + self.assertEqual( + self.example_dbtask[1], TaskHandler.get_task(self.session, "Task2") + ) def test_finding_task_by_id(self): clear_table(self.session, DBTask) From ef098486e10d82b72f437aa9c948fc29ebbbb0c5 Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Fri, 19 Jun 2026 11:41:29 +0800 Subject: [PATCH 06/10] fix: correct exception syntax --- app/core/exceptions.py | 10 +++++----- app/llm/base.py | 2 +- app/mcp/mcp_tools/conversion.py | 5 +++-- app/mcp/mcp_tools/miles_to_km.py | 11 ++++++----- app/routes/router_handler.py | 2 +- app/security/rate_limit.py | 4 +--- main.py | 2 +- tests/mcp/test_rate_limit.py | 11 ++--------- tests/mcp/test_root.py | 4 ++-- tests/mcp/test_validation_exception.py | 11 ++++++++++- 10 files changed, 32 insertions(+), 30 deletions(-) diff --git a/app/core/exceptions.py b/app/core/exceptions.py index 6fdf8af..b6b20ab 100644 --- a/app/core/exceptions.py +++ b/app/core/exceptions.py @@ -1,11 +1,11 @@ class ValidationError(Exception): - """Shared Validation exception for core business logic. - - Raised by core logic when input validation fails. - Fast API routes are responsible for catching this and + """Shared Validation exception for core business logic. + + Raised by core logic when input validation fails. + Fast API routes are responsible for catching this and converting it to an appropriate HTTP response. """ def __init__(self, message: str = "Validation error"): self.message = message - super().__init__(message) \ No newline at end of file + super().__init__(message) diff --git a/app/llm/base.py b/app/llm/base.py index b35798d..e45c5d9 100644 --- a/app/llm/base.py +++ b/app/llm/base.py @@ -41,7 +41,7 @@ class BaseLLMClient(Protocol): @property def provider_name(self) -> str: - """Returns the human-readable name of the LLM provider.""" + """Human-readable name of the LLM provider.""" ... async def generate(self, request: LLMRequest) -> LLMResponse: diff --git a/app/mcp/mcp_tools/conversion.py b/app/mcp/mcp_tools/conversion.py index fe13e84..8d7523e 100644 --- a/app/mcp/mcp_tools/conversion.py +++ b/app/mcp/mcp_tools/conversion.py @@ -1,4 +1,3 @@ - def miles_to_kilometers_converter(miles: float) -> float: """ Converter function that stores logic for miles to kilometers conversion. @@ -9,4 +8,6 @@ def miles_to_kilometers_converter(miles: float) -> float: Returns: float: Distance in kilometers. """ - return miles / 0.621371 # uses the same method what was in "miles_to_km" before it was moved here + return ( + miles / 0.621371 + ) # uses the same method what was in "miles_to_km" before it was moved here diff --git a/app/mcp/mcp_tools/miles_to_km.py b/app/mcp/mcp_tools/miles_to_km.py index e890a1f..910fcd8 100644 --- a/app/mcp/mcp_tools/miles_to_km.py +++ b/app/mcp/mcp_tools/miles_to_km.py @@ -2,8 +2,9 @@ import time from fastapi import APIRouter, HTTPException -from app.core.exceptions import ValidationError from pydantic import BaseModel, Field + +from app.core.exceptions import ValidationError from app.mcp.mcp_tools.conversion import miles_to_kilometers_converter router = APIRouter(prefix="", tags=["unit-conversion"]) @@ -56,7 +57,7 @@ def miles_to_kilometers_value(miles: float) -> float: if miles > MAX_TUTORIAL_MILES: raise ValidationError( "Distance is unrealistically large for this tutorial example." - ) + ) return miles_to_kilometers_converter(miles) @@ -82,10 +83,10 @@ def miles_to_kilometers( operation="miles_to_kilometers", audited_at=time.time(), ) - except ValidationError as exc: + except ValidationError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc - except ValueError as exc: - raise HTTPException(status_code=400, detail=str(exc)) from exc TOOL_DEFINITION = [ diff --git a/app/routes/router_handler.py b/app/routes/router_handler.py index 7024f8c..9765b47 100644 --- a/app/routes/router_handler.py +++ b/app/routes/router_handler.py @@ -27,7 +27,7 @@ def load_all_routes(cls): module = module_from_spec(spec) spec.loader.exec_module(module) - except UnicodeDecodeError, PermissionError: + except (UnicodeDecodeError, PermissionError): continue @classmethod diff --git a/app/security/rate_limit.py b/app/security/rate_limit.py index e09f6ca..8ca3cf4 100644 --- a/app/security/rate_limit.py +++ b/app/security/rate_limit.py @@ -3,15 +3,13 @@ from typing import override from fastapi import Request, Response -from starlette.status import HTTP_400_BAD_REQUEST, HTTP_429_TOO_MANY_REQUESTS from starlette.middleware.base import BaseHTTPMiddleware - +from starlette.status import HTTP_400_BAD_REQUEST, HTTP_429_TOO_MANY_REQUESTS # starlette is installed when fastapi is installed. # BaseHTTPMiddleware is the class we inherit from to create custom middleware in FastAPI/Starlette. - RATE_LIMIT_REQUESTS = 50 RATE_LIMIT_WINDOW_SECONDS = 3600 diff --git a/main.py b/main.py index 86945ba..253e06e 100644 --- a/main.py +++ b/main.py @@ -10,8 +10,8 @@ from app.mcp.mcp_resources.converter_resources import RESOURCE_DEFINITIONS from app.mcp.mcp_tools.miles_to_km import router as mile_to_km from app.routes.router_handler import Router -from app.utils.resource_utils import register_resources from app.security.rate_limit import RateLimitMiddleware +from app.utils.resource_utils import register_resources # FastAPI app for plain HTTP app = FastAPI( diff --git a/tests/mcp/test_rate_limit.py b/tests/mcp/test_rate_limit.py index ef92206..8b633cd 100644 --- a/tests/mcp/test_rate_limit.py +++ b/tests/mcp/test_rate_limit.py @@ -15,7 +15,6 @@ from app.security import rate_limit from app.security.rate_limit import RateLimitMiddleware - # Keep the test limit small so the tests can trigger the rate limiter quickly. TEST_RATE_LIMIT_REQUESTS = 2 TEST_RATE_LIMIT_WINDOW_SECONDS = 60 @@ -38,16 +37,10 @@ def test_rate_limit_blocks_after_limit(monkeypatch): # The real limit is higher, so monkeypatch lowers it for this test only. # Pytest restores the original values after the test finishes. - monkeypatch.setattr( - rate_limit, - "RATE_LIMIT_REQUESTS", - TEST_RATE_LIMIT_REQUESTS - ) + monkeypatch.setattr(rate_limit, "RATE_LIMIT_REQUESTS", TEST_RATE_LIMIT_REQUESTS) monkeypatch.setattr( - rate_limit, - "RATE_LIMIT_WINDOW_SECONDS", - TEST_RATE_LIMIT_WINDOW_SECONDS + rate_limit, "RATE_LIMIT_WINDOW_SECONDS", TEST_RATE_LIMIT_WINDOW_SECONDS ) # TestClient runs this small FastAPI app in the test process. diff --git a/tests/mcp/test_root.py b/tests/mcp/test_root.py index 95fb909..8bffb55 100644 --- a/tests/mcp/test_root.py +++ b/tests/mcp/test_root.py @@ -1,9 +1,9 @@ EXPECTED_STATUS_CODE = 200 # Expected status code to avoid magic number + def test_root(base_url, http_client): - """Verify root endpoint (/) works and returns results displaying available tools""" - + """Verify root endpoint (/) works and returns results displaying available tools.""" response = http_client.get(f"{base_url}/") assert response.status_code == EXPECTED_STATUS_CODE diff --git a/tests/mcp/test_validation_exception.py b/tests/mcp/test_validation_exception.py index d03a2d8..f9616d5 100644 --- a/tests/mcp/test_validation_exception.py +++ b/tests/mcp/test_validation_exception.py @@ -1,10 +1,12 @@ import pytest + from app.core.exceptions import ValidationError from app.mcp.mcp_tools.miles_to_km import miles_to_kilometers_value EXPECTED_KM = 1.609 CONVERSION_TOLERANCE = 0.001 + def test_known_conversion(): result = miles_to_kilometers_value(1) assert abs(result - EXPECTED_KM) < CONVERSION_TOLERANCE @@ -13,30 +15,37 @@ def test_known_conversion(): def test_validation_error_is_exception(): assert issubclass(ValidationError, Exception) + def test_negative_miles_raises_validation_error(): with pytest.raises(ValidationError): miles_to_kilometers_value(-1) + def test_zero_miles_raises_validation_error(): with pytest.raises(ValidationError): miles_to_kilometers_value(0) + def test_none_miles_raises_validation_error(): with pytest.raises(ValidationError): miles_to_kilometers_value(None) + def test_too_large_miles_raises_validation_error(): with pytest.raises(ValidationError): miles_to_kilometers_value(99999999) + def test_valid_miles_returns_float(): result = miles_to_kilometers_value(1) assert isinstance(result, float) + def test_http_exception_not_raised(): from fastapi import HTTPException + with pytest.raises(ValidationError): try: miles_to_kilometers_value(-1) except HTTPException: - pytest.fail("HTTPException should not be raised from core logic") \ No newline at end of file + pytest.fail("HTTPException should not be raised from core logic") From e15ec1bc5a8e537cb910095fafdc7b2dd9666127 Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Sun, 21 Jun 2026 21:54:44 +0800 Subject: [PATCH 07/10] fix: remove conflict markers from task_handler.py --- app/api/task_handler.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/app/api/task_handler.py b/app/api/task_handler.py index f1e59f3..44494ff 100644 --- a/app/api/task_handler.py +++ b/app/api/task_handler.py @@ -33,11 +33,7 @@ def create_task( @classmethod def add_task_to_db(cls, db: Session, task: PLTask): """Creates a task and adds it to the misc.""" -<<<<<<< HEAD # print(db) # debug print -======= - # print(db) # debug print ->>>>>>> upstream/main new_task_obj = DBTask( id=task.id or None, name=task.name, From 92670954761b58ab4ceb01577079d6f92baa761e Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Sun, 21 Jun 2026 21:58:09 +0800 Subject: [PATCH 08/10] fix: add return [] fallback in query_db to satisfy type checker --- app/services/database.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/app/services/database.py b/app/services/database.py index 3455fd9..b4acd7a 100644 --- a/app/services/database.py +++ b/app/services/database.py @@ -147,6 +147,7 @@ def query_db(db: Session, obj_ref, field: str, value) -> list[DBTask]: status_code=404, detail=f"Object with filed: {field}, does not exist", ) + return [] else: try: param = getattr(obj_ref, field) @@ -156,6 +157,7 @@ def query_db(db: Session, obj_ref, field: str, value) -> list[DBTask]: status_code=404, detail=f"Object with filed: {field}, does not exist", ) + return [] def get_session() -> Session: From d705bbf8b16df078715f0090fbea93cfc899c35d Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Sun, 21 Jun 2026 22:02:01 +0800 Subject: [PATCH 09/10] fix: resolve ty type checker errors --- app/api/task_handler.py | 6 +++--- app/llm/core/queue.py | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/app/api/task_handler.py b/app/api/task_handler.py index 44494ff..3dbe947 100644 --- a/app/api/task_handler.py +++ b/app/api/task_handler.py @@ -48,7 +48,7 @@ def add_task_to_db(cls, db: Session, task: PLTask): if not new_task_obj: raise Exception("No new task obj") - task.id = new_task_obj.id + task.id = int(new_task_obj.id) return task @@ -112,8 +112,8 @@ def complete_task(cls, db: Session, task: PLTask): if dbtask: try: - dbtask.completed = True - dbtask.task_ended = task.task_ended + dbtask.completed = True # type: ignore[assignment] + dbtask.task_ended = task.task_ended # type: ignore[assignment] db.commit() return task diff --git a/app/llm/core/queue.py b/app/llm/core/queue.py index c010256..5aaf4d6 100644 --- a/app/llm/core/queue.py +++ b/app/llm/core/queue.py @@ -144,7 +144,7 @@ async def wait_for_result( if isinstance(job, FinishedJob): outcome = job.outcome if outcome.status == "ok": - return cast(LLMResponse, outcome.root) # ty: ignore + return cast(LLMResponse, outcome.root) raise ValueError(f"Job failed: {outcome.root}") if anyio.current_time() > deadline: raise TimeoutError(f"Job {job_id} did not complete within {timeout}s") From 613e9391e0c4e6fd4cfd90030fb9db4309e895b6 Mon Sep 17 00:00:00 2001 From: 20149573 <20149573@tafe.wa.edu.au> Date: Sun, 21 Jun 2026 22:05:20 +0800 Subject: [PATCH 10/10] fix: lint error --- app/api/task_handler.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/app/api/task_handler.py b/app/api/task_handler.py index 3dbe947..fcfb30d 100644 --- a/app/api/task_handler.py +++ b/app/api/task_handler.py @@ -48,7 +48,7 @@ def add_task_to_db(cls, db: Session, task: PLTask): if not new_task_obj: raise Exception("No new task obj") - task.id = int(new_task_obj.id) + task.id = new_task_obj.id # ty: ignore return task @@ -112,8 +112,8 @@ def complete_task(cls, db: Session, task: PLTask): if dbtask: try: - dbtask.completed = True # type: ignore[assignment] - dbtask.task_ended = task.task_ended # type: ignore[assignment] + dbtask.completed = True # ty: ignore + dbtask.task_ended = task.task_ended # ty: ignore db.commit() return task