main.py 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195
  1. from typing import Any
  2. from fastapi import Depends, FastAPI, Header, HTTPException, Query
  3. from fastapi.middleware.cors import CORSMiddleware
  4. from pydantic import BaseModel, Field
  5. from .auth import create_token, revoke_token, validate_token, verify_credentials
  6. from .config import settings
  7. from .db import ensure_annotation_table
  8. from .services.data_service import data_service
  9. class TimePointInput(BaseModel):
  10. index: int = Field(ge=0)
  11. sampleTime: str
  12. files: dict[str, Any]
  13. class WaveWindowInput(BaseModel):
  14. pointName: str
  15. measurementTypes: list[str]
  16. points: list[TimePointInput]
  17. maxPoints: int = Field(default=200000, ge=256, le=200000)
  18. noSampling: bool = False
  19. class AnnotationInput(BaseModel):
  20. wave_file_id: int = Field(gt=0)
  21. label: str = Field(min_length=1, max_length=8)
  22. period_start: int = Field(gt=0)
  23. period_end: int = Field(gt=0)
  24. sample_index_start: int = Field(ge=0)
  25. sample_index_end: int = Field(ge=0)
  26. class AnnotationQueryInput(BaseModel):
  27. wave_file_ids: list[int] = Field(default=[])
  28. class LoginInput(BaseModel):
  29. username: str = Field(min_length=1)
  30. password: str = Field(min_length=1)
  31. app = FastAPI(
  32. title="压缩机故障预测波形服务",
  33. version="0.1.0",
  34. )
  35. app.add_middleware(
  36. CORSMiddleware,
  37. allow_origins=[settings.cors_origin, "http://127.0.0.1:5173"],
  38. allow_credentials=True,
  39. allow_methods=["*"],
  40. allow_headers=["*"],
  41. )
  42. def require_auth(authorization: str | None = Header(default=None)) -> str:
  43. token = (authorization or "").removeprefix("Bearer ").strip()
  44. if not validate_token(token):
  45. raise HTTPException(status_code=401, detail="未登录或登录已过期")
  46. return token
  47. @app.on_event("startup")
  48. def startup() -> None:
  49. try:
  50. ensure_annotation_table()
  51. except Exception:
  52. # 数据库不可达时跳过建表,演示模式仍可正常启动。
  53. pass
  54. @app.post("/api/login")
  55. def login(payload: LoginInput) -> dict[str, Any]:
  56. if not verify_credentials(payload.username, payload.password):
  57. raise HTTPException(status_code=401, detail="用户名或密码错误")
  58. return {"token": create_token(), "username": payload.username}
  59. @app.post("/api/logout")
  60. def logout(authorization: str | None = Header(default=None)) -> dict[str, Any]:
  61. token = (authorization or "").removeprefix("Bearer ").strip()
  62. revoke_token(token)
  63. return {"ok": True}
  64. @app.get("/api/health")
  65. def health(_auth: str = Depends(require_auth)) -> dict[str, Any]:
  66. return data_service.health()
  67. @app.get("/api/query-options")
  68. def query_options(_auth: str = Depends(require_auth)) -> dict[str, Any]:
  69. try:
  70. return data_service.query_options()
  71. except Exception as error:
  72. raise HTTPException(status_code=503, detail=str(error)) from error
  73. @app.get("/api/time-points")
  74. def time_points(
  75. point_name: str = Query(min_length=1),
  76. measurement_types: list[str] = Query(default=[]),
  77. min_time: str | None = None,
  78. max_time: str | None = None,
  79. _auth: str = Depends(require_auth),
  80. ) -> dict[str, Any]:
  81. try:
  82. return data_service.time_points(point_name, measurement_types, min_time, max_time)
  83. except ValueError as error:
  84. raise HTTPException(status_code=400, detail=str(error)) from error
  85. except Exception as error:
  86. raise HTTPException(status_code=503, detail=str(error)) from error
  87. @app.post("/api/wave-window")
  88. def wave_window(payload: WaveWindowInput, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  89. try:
  90. return data_service.wave_window(
  91. payload.pointName,
  92. payload.measurementTypes,
  93. [point.model_dump() if hasattr(point, "model_dump") else point.dict() for point in payload.points],
  94. payload.maxPoints,
  95. payload.noSampling,
  96. )
  97. except ValueError as error:
  98. raise HTTPException(status_code=400, detail=str(error)) from error
  99. except Exception as error:
  100. raise HTTPException(status_code=503, detail=str(error)) from error
  101. @app.get("/api/wave-files/{wave_file_id}/periods/{period_number}")
  102. def period_detail(
  103. wave_file_id: int,
  104. period_number: int,
  105. _auth: str = Depends(require_auth),
  106. ) -> dict[str, Any]:
  107. try:
  108. return data_service.period_detail(wave_file_id, period_number)
  109. except ValueError as error:
  110. raise HTTPException(status_code=400, detail=str(error)) from error
  111. except Exception as error:
  112. raise HTTPException(status_code=503, detail=str(error)) from error
  113. @app.get("/api/annotation-config")
  114. def annotation_config(_auth: str = Depends(require_auth)) -> dict[str, Any]:
  115. return data_service.annotation_config()
  116. @app.get("/api/annotations")
  117. def list_annotations(
  118. wave_file_ids: list[int] = Query(default=[]),
  119. _auth: str = Depends(require_auth),
  120. ) -> dict[str, Any]:
  121. try:
  122. return data_service.list_annotations(wave_file_ids)
  123. except ValueError as error:
  124. raise HTTPException(status_code=400, detail=str(error)) from error
  125. except Exception as error:
  126. raise HTTPException(status_code=503, detail=str(error)) from error
  127. @app.post("/api/annotations/query")
  128. def query_annotations(
  129. payload: AnnotationQueryInput,
  130. _auth: str = Depends(require_auth),
  131. ) -> dict[str, Any]:
  132. try:
  133. return data_service.list_annotations(payload.wave_file_ids)
  134. except ValueError as error:
  135. raise HTTPException(status_code=400, detail=str(error)) from error
  136. except Exception as error:
  137. raise HTTPException(status_code=503, detail=str(error)) from error
  138. @app.post("/api/annotations")
  139. def create_annotation(payload: AnnotationInput, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  140. try:
  141. return data_service.create_annotation(payload.model_dump())
  142. except ValueError as error:
  143. raise HTTPException(status_code=400, detail=str(error)) from error
  144. except Exception as error:
  145. raise HTTPException(status_code=503, detail=str(error)) from error
  146. @app.delete("/api/annotations/{annotation_id}")
  147. def delete_annotation(annotation_id: int, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  148. try:
  149. return data_service.delete_annotation(annotation_id)
  150. except ValueError as error:
  151. raise HTTPException(status_code=404, detail=str(error)) from error
  152. except Exception as error:
  153. raise HTTPException(status_code=503, detail=str(error)) from error