File size: 3,092 Bytes
4a93e6c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e2a227a
 
 
 
 
 
 
 
 
 
4a93e6c
 
 
e2a227a
 
4a93e6c
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
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)