Repository navigation
Expand file tree
/
Copy pathapp.py
More file actions
155 lines (122 loc) · 4.82 KB
/
Copy pathapp.py
File metadata and controls
155 lines (122 loc) · 4.82 KB
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
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
import os
from flask import Flask, request, jsonify, g
from flask_cors import CORS
import io
from PIL import Image
import numpy as np
import firebase_admin
import datetime
from model_load import model, class_names
from google.cloud import storage, firestore
from auth_middleware import firebase_authentication_middleware
from flask_swagger_ui import get_swaggerui_blueprint
os.environ['TF_ENABLE_ONEDNN_OPTS'] = '0'
app = Flask(__name__)
CORS(app, resources={r"/*": {"origins": "*"}})
prod = os.environ.get('PRODUCTION', "False").lower() == "true"
cred = None
if prod:
cred = firebase_admin.credentials.ApplicationDefault()
else:
cred = firebase_admin.credentials.Certificate('venv/serviceAccount.json')
firebase_admin.initialize_app(cred)
SWAGGER_URL = '/api-docs'
API_URL = '/static/openapi.json'
swaggerui_blueprint = get_swaggerui_blueprint(
SWAGGER_URL,
API_URL,
config={
'E-iqro': "API Docs"
}
)
app.register_blueprint(swaggerui_blueprint)
def upload_image_to_gcs(bucket_name, file_stream, destination_blob_name):
if prod:
storage_client = storage.Client()
else:
storage_client = storage.Client.from_service_account_json('venv/serviceAccount.json')
bucket = storage_client.bucket(bucket_name)
blob = bucket.blob(destination_blob_name)
blob.upload_from_file(file_stream, content_type='image/jpeg')
url = blob.public_url
print(f"File uploaded to {url}.")
return url
def save_prediction_to_firestore(predicted_class, confidence, image_url, user_id):
if prod:
db = firestore.Client()
else:
db = firestore.Client.from_service_account_json('venv/serviceAccount.json')
new_prediction_ref = db.collection('history').document()
data = {
'uid': user_id,
'predicted_class': predicted_class,
'confidence': confidence,
'image_url': image_url,
'timestamp': datetime.datetime.now()
}
new_prediction_ref.set(data)
print(f"Prediction saved to Firestore with image URL: {image_url}")
def preprocess_image_as_array(image):
im = Image.open(image).convert('RGB')
im = im.resize((224, 224))
return np.asarray(im)
def predict_image_class(model, img_array, class_names, threshold=0.7):
img_batch = np.expand_dims(img_array, axis=0)
predictions = model.predict(img_batch)[0]
predicted_class_index = np.argmax(predictions)
predicted_class_score = predictions[predicted_class_index]
predicted_class = class_names[predicted_class_index]
if predicted_class_score >= threshold:
result = {
"predicted": predicted_class,
"confidence": float(predicted_class_score),
}
return result
else:
return None
@app.route('/v1/predict', methods=['POST'])
@firebase_authentication_middleware
def predict():
if 'image' not in request.files:
return jsonify({'error': 'No file part'})
file = request.files['image']
if file.filename == '':
return jsonify({'error': 'No selected file'})
try:
img_array = preprocess_image_as_array(file)
predicted_class = predict_image_class(model, img_array, class_names, threshold=0.7)
if predicted_class:
bucket_name = 'images_from_predict'
destination_blob_name = "image_predict_" + datetime.datetime.now().strftime("%Y%m%d%H%M%S") + "_" + file.filename
file.seek(0)
file_stream = io.BytesIO(file.read())
file_stream.seek(0)
image_url = upload_image_to_gcs(bucket_name, file_stream, destination_blob_name)
user_id = g.uid
save_prediction_to_firestore(predicted_class, predicted_class['confidence'], image_url, user_id)
return jsonify({'result': predicted_class['predicted'], 'confidence': predicted_class['confidence'],'uid' : user_id, 'image_url' : image_url})
else:
return jsonify({'error': 'prediction invalid'})
except Exception as e:
return jsonify({'error': str(e)})
@app.route('/v1/history', methods=['GET'])
@firebase_authentication_middleware
def get_history():
try:
user_id = g.uid
if not user_id:
return jsonify({'error': 'No user ID provided'}), 400
if prod:
db = firestore.Client()
else:
db = firestore.Client.from_service_account_json('venv/serviceAccount.json')
history_ref = db.collection('history').where('uid', '==', user_id)
history_docs = history_ref.stream()
history_list = []
for doc in history_docs:
history_list.append(doc.to_dict())
return jsonify({'history': history_list}), 200
except Exception as e:
return jsonify({'error': str(e)}), 500
if __name__ == '__main__':
app.run(host='0.0.0.0', port=8080, debug=False)