AhmedMahmoud165's picture
Update app.py
e2a227a verified
Raw
History Blame Contribute Delete
3.09 kB
from gradio_client import Client
import os
import time
from concurrent.futures import ThreadPoolExecutor, TimeoutError
from dotenv import load_dotenv
load_dotenv()
token = os.getenv("HF_TOKEN")
HF_SPACE_URL = "AhmedMahmoud165/Machine-learning"
_client_instance = None
_circuit_open_until = 0.0
_consecutive_failures = 0
_executor = ThreadPoolExecutor(max_workers=4)
def _get_client():
global _client_instance
if _client_instance is None:
try:
kwargs = {}
if token:
kwargs["hf_token"] = token
_client_instance = Client(HF_SPACE_URL, **kwargs)
except Exception:
_client_instance = None
return _client_instance
def local_fallback_predict(payload_data):
src_bytes = payload_data.get('src_bytes', 0)
count = payload_data.get('count', 1)
duration = payload_data.get('duration', 0.0)
protocol_type = str(payload_data.get('protocol_type', '')).upper()
if src_bytes > 50000 or count > 100 or (duration > 0 and (src_bytes / max(duration, 0.001)) > 100000):
return 1
if protocol_type == "UDP" and src_bytes > 1500:
return 1
return 0
def analyze_payload_ML(payload_data):
global _circuit_open_until, _consecutive_failures, _client_instance
now = time.time()
if now < _circuit_open_until:
return local_fallback_predict(payload_data)
def _predict_task():
client = _get_client()
if client is None:
raise RuntimeError("Could not initialize Gradio client")
duration = payload_data.get('duration', 0.0)
protocol_type = payload_data.get('protocol_type', 'TCP')
service = payload_data.get('service', 'http')
src_bytes = payload_data.get('src_bytes', 0)
dst_bytes = payload_data.get('dst_bytes', 0)
count = payload_data.get('count', 1)
result = client.predict(
duration,
protocol_type,
service,
src_bytes,
dst_bytes,
count,
api_name="/predict"
)
# معالجة صحيحة لشكل البيانات الراجع من Gradio Client (سواء List أو Dict أو قيمة مباشرة)
if isinstance(result, (list, tuple)):
prediction = result[0]
elif isinstance(result, dict):
prediction = result.get('prediction', 0)
else:
prediction = result
return int(prediction)
try:
future = _executor.submit(_predict_task)
# رفع الـ Timeout إلى 10.0 ثواني لضمان استقرار الطلبات
prediction = future.result(timeout=10.0)
_consecutive_failures = 0
return prediction
except Exception as e:
_consecutive_failures += 1
_client_instance = None
if _consecutive_failures >= 3:
_circuit_open_until = now + 60.0
print(f"ML Circuit breaker opened until {_circuit_open_until} due to error: {e}")
return local_fallback_predict(payload_data)