main.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378
  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. from .services.pretrain_service import pretrain_service
  10. class TimePointInput(BaseModel):
  11. index: int = Field(ge=0)
  12. sampleTime: str
  13. files: dict[str, Any]
  14. siteValues: dict[str, float | None] = Field(default={})
  15. class WaveWindowInput(BaseModel):
  16. devicePart: str = Field(min_length=1)
  17. devicePoints: list[str]
  18. points: list[TimePointInput]
  19. maxPoints: int = Field(default=200000, ge=256, le=200000)
  20. noSampling: bool = False
  21. firstCycleOnly: bool = False
  22. sitePoints: list[str] = Field(default=[])
  23. class AnnotationInput(BaseModel):
  24. wave_file_id: int = Field(gt=0)
  25. label: str = Field(min_length=1, max_length=8)
  26. period_start: int = Field(gt=0)
  27. period_end: int = Field(gt=0)
  28. sample_index_start: int = Field(ge=0)
  29. sample_index_end: int = Field(ge=0)
  30. class AnnotationQueryInput(BaseModel):
  31. wave_file_ids: list[int] = Field(default=[])
  32. class AlarmScanInput(BaseModel):
  33. forecast_time: str = Field(min_length=1)
  34. device_code: str = Field(min_length=1)
  35. alarm_type: str = Field(min_length=1)
  36. hours: float = Field(gt=0)
  37. class PressureAngleSyncInput(BaseModel):
  38. device_part: str = Field(min_length=1)
  39. device_points: list[str] = Field(default=[])
  40. class PressureAngleQueryInput(BaseModel):
  41. device_part: str = Field(min_length=1)
  42. device_points: list[str] = Field(default=[])
  43. min_time: str | None = None
  44. max_time: str | None = None
  45. angles: list[int] = Field(default=[])
  46. class PretrainRecordInput(BaseModel):
  47. id: int = Field(gt=0)
  48. point_name: str = Field(min_length=1)
  49. measurement_type: str = Field(min_length=1)
  50. rpm: float | None = None
  51. sample_time: str = Field(min_length=1)
  52. file_name: str = Field(default="")
  53. class LoginInput(BaseModel):
  54. username: str = Field(min_length=1)
  55. password: str = Field(min_length=1)
  56. app = FastAPI(
  57. title="压缩机故障预测波形服务",
  58. version="0.1.0",
  59. )
  60. app.add_middleware(
  61. CORSMiddleware,
  62. allow_origins=[settings.cors_origin, "http://127.0.0.1:5173"],
  63. allow_credentials=True,
  64. allow_methods=["*"],
  65. allow_headers=["*"],
  66. )
  67. def require_auth(authorization: str | None = Header(default=None)) -> str:
  68. token = (authorization or "").removeprefix("Bearer ").strip()
  69. if not validate_token(token):
  70. raise HTTPException(status_code=401, detail="未登录或登录已过期")
  71. return token
  72. @app.on_event("startup")
  73. def startup() -> None:
  74. try:
  75. ensure_annotation_table()
  76. except Exception:
  77. # 数据库不可达时跳过建表,演示模式仍可正常启动。
  78. pass
  79. @app.post("/api/login")
  80. def login(payload: LoginInput) -> dict[str, Any]:
  81. if not verify_credentials(payload.username, payload.password):
  82. raise HTTPException(status_code=401, detail="用户名或密码错误")
  83. return {"token": create_token(), "username": payload.username}
  84. @app.post("/api/logout")
  85. def logout(authorization: str | None = Header(default=None)) -> dict[str, Any]:
  86. token = (authorization or "").removeprefix("Bearer ").strip()
  87. revoke_token(token)
  88. return {"ok": True}
  89. @app.get("/api/health")
  90. def health(_auth: str = Depends(require_auth)) -> dict[str, Any]:
  91. return data_service.health()
  92. @app.get("/api/query-options")
  93. def query_options(_auth: str = Depends(require_auth)) -> dict[str, Any]:
  94. try:
  95. result = data_service.query_options()
  96. # Counts are keyed by the full point_name so the UI can scope them to
  97. # the currently selected device_part.
  98. result["pretrainCounts"] = pretrain_service.recorded_counts()
  99. result["abnormalCounts"] = data_service.abnormal_counts()
  100. return result
  101. except Exception as error:
  102. raise HTTPException(status_code=503, detail=str(error)) from error
  103. @app.get("/api/time-points")
  104. def time_points(
  105. device_part: str = Query(min_length=1),
  106. device_points: list[str] = Query(default=[]),
  107. min_time: str | None = None,
  108. max_time: str | None = None,
  109. include_stopped: bool = False,
  110. min_status: int | None = Query(default=None, ge=0),
  111. status_filter: list[str] = Query(default=[]),
  112. site_points: list[str] = Query(default=[]),
  113. _auth: str = Depends(require_auth),
  114. ) -> dict[str, Any]:
  115. try:
  116. return data_service.time_points(
  117. device_part,
  118. device_points,
  119. min_time,
  120. max_time,
  121. include_stopped,
  122. min_status,
  123. status_filter,
  124. site_points,
  125. )
  126. except ValueError as error:
  127. raise HTTPException(status_code=400, detail=str(error)) from error
  128. except Exception as error:
  129. raise HTTPException(status_code=503, detail=str(error)) from error
  130. @app.get("/api/site-points")
  131. def site_points(
  132. device_part: str = Query(min_length=1),
  133. _auth: str = Depends(require_auth),
  134. ) -> dict[str, Any]:
  135. try:
  136. return data_service.site_points(device_part)
  137. except Exception as error:
  138. raise HTTPException(status_code=503, detail=str(error)) from error
  139. @app.get("/api/faults")
  140. def faults(
  141. device_part: str = Query(min_length=1),
  142. min_time: str | None = None,
  143. max_time: str | None = None,
  144. _auth: str = Depends(require_auth),
  145. ) -> dict[str, Any]:
  146. try:
  147. return data_service.faults(device_part, min_time, max_time)
  148. except ValueError as error:
  149. raise HTTPException(status_code=400, detail=str(error)) from error
  150. except Exception as error:
  151. raise HTTPException(status_code=503, detail=str(error)) from error
  152. @app.get("/api/tspluse-ruler")
  153. def tspluse_ruler(_auth: str = Depends(require_auth)) -> dict[str, Any]:
  154. try:
  155. return data_service.tspluse_ruler()
  156. except Exception as error:
  157. raise HTTPException(status_code=503, detail=str(error)) from error
  158. @app.get("/api/alarms")
  159. def list_alarms(
  160. current_time: str | None = None,
  161. _auth: str = Depends(require_auth),
  162. ) -> dict[str, Any]:
  163. try:
  164. return data_service.list_alarms(current_time)
  165. except ValueError as error:
  166. raise HTTPException(status_code=400, detail=str(error)) from error
  167. except Exception as error:
  168. raise HTTPException(status_code=503, detail=str(error)) from error
  169. @app.get("/api/alarm-compressor-options")
  170. def alarm_compressor_options(_auth: str = Depends(require_auth)) -> dict[str, Any]:
  171. try:
  172. return data_service.alarm_compressor_options()
  173. except Exception as error:
  174. raise HTTPException(status_code=503, detail=str(error)) from error
  175. @app.post("/api/alarm-scan")
  176. def alarm_scan(payload: AlarmScanInput, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  177. try:
  178. return data_service.scan_alarm(payload.model_dump())
  179. except ValueError as error:
  180. raise HTTPException(status_code=400, detail=str(error)) from error
  181. except Exception as error:
  182. raise HTTPException(status_code=503, detail=str(error)) from error
  183. @app.post("/api/pressure-angle/sync")
  184. def pressure_angle_sync(payload: PressureAngleSyncInput, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  185. try:
  186. return data_service.sync_pressure_angle(payload.device_part, payload.device_points)
  187. except ValueError as error:
  188. raise HTTPException(status_code=400, detail=str(error)) from error
  189. except Exception as error:
  190. raise HTTPException(status_code=503, detail=str(error)) from error
  191. @app.post("/api/pressure-angle/query")
  192. def pressure_angle_query(payload: PressureAngleQueryInput, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  193. try:
  194. return data_service.query_pressure_angle(
  195. payload.device_part,
  196. payload.device_points,
  197. payload.min_time,
  198. payload.max_time,
  199. payload.angles,
  200. )
  201. except ValueError as error:
  202. raise HTTPException(status_code=400, detail=str(error)) from error
  203. except Exception as error:
  204. raise HTTPException(status_code=503, detail=str(error)) from error
  205. @app.post("/api/wave-window")
  206. def wave_window(payload: WaveWindowInput, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  207. try:
  208. return data_service.wave_window(
  209. payload.devicePart,
  210. payload.devicePoints,
  211. [point.model_dump() if hasattr(point, "model_dump") else point.dict() for point in payload.points],
  212. payload.maxPoints,
  213. payload.noSampling,
  214. payload.firstCycleOnly,
  215. payload.sitePoints,
  216. )
  217. except ValueError as error:
  218. raise HTTPException(status_code=400, detail=str(error)) from error
  219. except Exception as error:
  220. raise HTTPException(status_code=503, detail=str(error)) from error
  221. @app.get("/api/wave-files/{wave_file_id}/periods/{period_number}")
  222. def period_detail(
  223. wave_file_id: int,
  224. period_number: int,
  225. _auth: str = Depends(require_auth),
  226. ) -> dict[str, Any]:
  227. try:
  228. return data_service.period_detail(wave_file_id, period_number)
  229. except ValueError as error:
  230. raise HTTPException(status_code=400, detail=str(error)) from error
  231. except Exception as error:
  232. raise HTTPException(status_code=503, detail=str(error)) from error
  233. @app.get("/api/annotation-config")
  234. def annotation_config(_auth: str = Depends(require_auth)) -> dict[str, Any]:
  235. return data_service.annotation_config()
  236. @app.get("/api/annotations")
  237. def list_annotations(
  238. wave_file_ids: list[int] = Query(default=[]),
  239. _auth: str = Depends(require_auth),
  240. ) -> dict[str, Any]:
  241. try:
  242. return data_service.list_annotations(wave_file_ids)
  243. except ValueError as error:
  244. raise HTTPException(status_code=400, detail=str(error)) from error
  245. except Exception as error:
  246. raise HTTPException(status_code=503, detail=str(error)) from error
  247. @app.post("/api/annotations/query")
  248. def query_annotations(
  249. payload: AnnotationQueryInput,
  250. _auth: str = Depends(require_auth),
  251. ) -> dict[str, Any]:
  252. try:
  253. return data_service.list_annotations(payload.wave_file_ids)
  254. except ValueError as error:
  255. raise HTTPException(status_code=400, detail=str(error)) from error
  256. except Exception as error:
  257. raise HTTPException(status_code=503, detail=str(error)) from error
  258. @app.post("/api/annotations")
  259. def create_annotation(payload: AnnotationInput, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  260. try:
  261. return data_service.create_annotation(payload.model_dump())
  262. except ValueError as error:
  263. raise HTTPException(status_code=400, detail=str(error)) from error
  264. except Exception as error:
  265. raise HTTPException(status_code=503, detail=str(error)) from error
  266. @app.delete("/api/annotations/{annotation_id}")
  267. def delete_annotation(annotation_id: int, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  268. try:
  269. return data_service.delete_annotation(annotation_id)
  270. except ValueError as error:
  271. raise HTTPException(status_code=404, detail=str(error)) from error
  272. except Exception as error:
  273. raise HTTPException(status_code=503, detail=str(error)) from error
  274. @app.get("/api/pretrain-records")
  275. def pretrain_records_status(
  276. ids: list[int] = Query(default=[]),
  277. _auth: str = Depends(require_auth),
  278. ) -> dict[str, Any]:
  279. try:
  280. return {"recorded": pretrain_service.recorded_ids(ids)}
  281. except Exception as error:
  282. raise HTTPException(status_code=503, detail=str(error)) from error
  283. @app.post("/api/pretrain-records")
  284. def create_pretrain_record(
  285. payload: PretrainRecordInput,
  286. _auth: str = Depends(require_auth),
  287. ) -> dict[str, Any]:
  288. try:
  289. recorded, already = pretrain_service.record(payload.model_dump())
  290. except ValueError as error:
  291. raise HTTPException(status_code=400, detail=str(error)) from error
  292. except Exception as error:
  293. raise HTTPException(status_code=503, detail=str(error)) from error
  294. return {"recorded": recorded, "already": already}
  295. @app.delete("/api/pretrain-records/{file_id}")
  296. def delete_pretrain_record(file_id: int, _auth: str = Depends(require_auth)) -> dict[str, Any]:
  297. try:
  298. deleted = pretrain_service.remove(file_id)
  299. except ValueError as error:
  300. raise HTTPException(status_code=400, detail=str(error)) from error
  301. except Exception as error:
  302. raise HTTPException(status_code=503, detail=str(error)) from error
  303. if not deleted:
  304. raise HTTPException(status_code=404, detail="CSV 中不存在该记录")
  305. return {"deleted": deleted}