-
-
Notifications
You must be signed in to change notification settings - Fork 38.2k
Add facebox auth #15439
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Add facebox auth #15439
Changes from 9 commits
e50ce87
61105e4
e6cf59e
13f70cd
53a8846
ffe1c5c
259bb13
a5d981b
c2fa465
788a8e5
ccb7247
a00389b
6aaa88a
e07af2e
08d66c9
42529d7
c0298b6
72f08fd
df5e05d
509a4b5
2cfc529
7006994
6154103
f673150
136e7c8
303ff24
871c272
ad5008c
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -17,25 +17,29 @@ | |
| from homeassistant.components.image_processing import ( | ||
| PLATFORM_SCHEMA, ImageProcessingFaceEntity, ATTR_CONFIDENCE, CONF_SOURCE, | ||
| CONF_ENTITY_ID, CONF_NAME, DOMAIN) | ||
| from homeassistant.const import (CONF_IP_ADDRESS, CONF_PORT) | ||
| from homeassistant.const import ( | ||
| CONF_IP_ADDRESS, CONF_PORT, CONF_PASSWORD, CONF_USERNAME, | ||
| HTTP_BAD_REQUEST, HTTP_UNAUTHORIZED) | ||
|
|
||
| _LOGGER = logging.getLogger(__name__) | ||
|
|
||
| ATTR_BOUNDING_BOX = 'bounding_box' | ||
| ATTR_CLASSIFIER = 'classifier' | ||
|
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Really I see this in |
||
| ATTR_IMAGE_ID = 'image_id' | ||
| ATTR_ID = 'id' | ||
| ATTR_MATCHED = 'matched' | ||
| FACEBOX_NAME = 'name' | ||
| CLASSIFIER = 'facebox' | ||
| DATA_FACEBOX = 'facebox_classifiers' | ||
| EVENT_CLASSIFIER_TEACH = 'image_processing.teach_classifier' | ||
| FILE_PATH = 'file_path' | ||
| SERVICE_TEACH_FACE = 'facebox_teach_face' | ||
| TIMEOUT = 9 | ||
|
|
||
|
|
||
| PLATFORM_SCHEMA = PLATFORM_SCHEMA.extend({ | ||
| vol.Required(CONF_IP_ADDRESS): cv.string, | ||
| vol.Required(CONF_PORT): cv.port, | ||
| vol.Optional(CONF_USERNAME): cv.string, | ||
| vol.Optional(CONF_PASSWORD): cv.string, | ||
| }) | ||
|
|
||
| SERVICE_TEACH_SCHEMA = vol.Schema({ | ||
|
|
@@ -63,10 +67,10 @@ def parse_faces(api_faces): | |
| for entry in api_faces: | ||
| face = {} | ||
| if entry['matched']: # This data is only in matched faces. | ||
| face[ATTR_NAME] = entry['name'] | ||
| face[FACEBOX_NAME] = entry['name'] | ||
| face[ATTR_IMAGE_ID] = entry['id'] | ||
| else: # Lets be explicit. | ||
| face[ATTR_NAME] = None | ||
| face[FACEBOX_NAME] = None | ||
| face[ATTR_IMAGE_ID] = None | ||
| face[ATTR_CONFIDENCE] = round(100.0*entry['confidence'], 2) | ||
| face[ATTR_MATCHED] = entry['matched'] | ||
|
|
@@ -75,15 +79,38 @@ def parse_faces(api_faces): | |
| return known_faces | ||
|
|
||
|
|
||
| def post_image(url, image): | ||
| def post_image(url, username, password, image): | ||
| """Post an image to the classifier.""" | ||
| try: | ||
| response = requests.post( | ||
| url, | ||
| auth=requests.auth.HTTPBasicAuth(username, password), | ||
| json={"base64": encode_image(image)}, | ||
| timeout=TIMEOUT | ||
| ) | ||
| return response | ||
| if response.status_code == HTTP_UNAUTHORIZED: | ||
| _LOGGER.error("AuthenticationError on %s", CLASSIFIER) | ||
| return None | ||
| else: | ||
| return response | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Since we return something here, the other places that this function can exit at are lacking a return statement. If a function returns something, all exits need to return something. Consistency is good.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. OK |
||
| except requests.exceptions.ConnectionError: | ||
| _LOGGER.error("ConnectionError: Is %s running?", CLASSIFIER) | ||
|
|
||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Same here.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. ok |
||
|
|
||
| def teach_file(url, username, password, name, file_path): | ||
| """Teach the classifier a name associated with a file.""" | ||
| try: | ||
| with open(file_path, 'rb') as open_file: | ||
| response = requests.post( | ||
| url, | ||
| auth=requests.auth.HTTPBasicAuth(username, password), | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. What if username or password is None?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. They are by default and its OK
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Don't send in auth if username and password are kwargs = {}
if username:
kwargs['auth'] = requests.auth.HTTPBasicAuth(username, password)
response = request.post(url, **kwargs) |
||
| data={FACEBOX_NAME: name, ATTR_ID: file_path}, | ||
| files={'file': open_file}, | ||
| ) | ||
| if response.status_code == HTTP_UNAUTHORIZED: | ||
| _LOGGER.error("AuthenticationError on %s", CLASSIFIER) | ||
| elif response.status_code == HTTP_BAD_REQUEST: | ||
| _LOGGER.error("%s teaching of file %s failed with message:%s", | ||
| CLASSIFIER, file_path, response.text) | ||
| except requests.exceptions.ConnectionError: | ||
| _LOGGER.error("ConnectionError: Is %s running?", CLASSIFIER) | ||
|
|
||
|
|
@@ -104,12 +131,15 @@ def setup_platform(hass, config, add_devices, discovery_info=None): | |
| if DATA_FACEBOX not in hass.data: | ||
| hass.data[DATA_FACEBOX] = [] | ||
|
|
||
| ip = config[CONF_IP_ADDRESS] | ||
| port = config[CONF_PORT] | ||
| username = config.get(CONF_USERNAME) | ||
| password = config.get(CONF_PASSWORD) | ||
|
|
||
| entities = [] | ||
| for camera in config[CONF_SOURCE]: | ||
| facebox = FaceClassifyEntity( | ||
| config[CONF_IP_ADDRESS], | ||
| config[CONF_PORT], | ||
| camera[CONF_ENTITY_ID], | ||
| ip, port, username, password, camera[CONF_ENTITY_ID], | ||
| camera.get(CONF_NAME)) | ||
| entities.append(facebox) | ||
| hass.data[DATA_FACEBOX].append(facebox) | ||
|
|
@@ -129,33 +159,33 @@ def service_handle(service): | |
| classifier.teach(name, file_path) | ||
|
|
||
| hass.services.register( | ||
| DOMAIN, | ||
| SERVICE_TEACH_FACE, | ||
| service_handle, | ||
| DOMAIN, SERVICE_TEACH_FACE, service_handle, | ||
| schema=SERVICE_TEACH_SCHEMA) | ||
|
|
||
|
|
||
| class FaceClassifyEntity(ImageProcessingFaceEntity): | ||
| """Perform a face classification.""" | ||
|
|
||
| def __init__(self, ip, port, camera_entity, name=None): | ||
| def __init__(self, ip, port, username, password, camera_entity, name=None): | ||
| """Init with the API key and model id.""" | ||
| super().__init__() | ||
| self._url_check = "http://{}:{}/{}/check".format(ip, port, CLASSIFIER) | ||
| self._url_teach = "http://{}:{}/{}/teach".format(ip, port, CLASSIFIER) | ||
| self._username = username | ||
| self._password = password | ||
| self._camera = camera_entity | ||
| if name: | ||
| self._name = name | ||
| else: | ||
| camera_name = split_entity_id(camera_entity)[1] | ||
| self._name = "{} {}".format( | ||
| CLASSIFIER, camera_name) | ||
| self._name = "{} {}".format(CLASSIFIER, camera_name) | ||
| self._matched = {} | ||
|
|
||
| def process_image(self, image): | ||
| """Process an image.""" | ||
| response = post_image(self._url_check, image) | ||
| if response is not None: | ||
| response = post_image( | ||
| self._url_check, self._username, self._password, image) | ||
| if response: | ||
| response_json = response.json() | ||
| if response_json['success']: | ||
| total_faces = response_json['facesCount'] | ||
|
|
@@ -173,34 +203,8 @@ def teach(self, name, file_path): | |
| if (not self.hass.config.is_allowed_path(file_path) | ||
| or not valid_file_path(file_path)): | ||
| return | ||
| with open(file_path, 'rb') as open_file: | ||
| response = requests.post( | ||
| self._url_teach, | ||
| data={ATTR_NAME: name, 'id': file_path}, | ||
| files={'file': open_file}) | ||
|
|
||
| if response.status_code == 200: | ||
| self.hass.bus.fire( | ||
| EVENT_CLASSIFIER_TEACH, { | ||
| ATTR_CLASSIFIER: CLASSIFIER, | ||
| ATTR_NAME: name, | ||
| FILE_PATH: file_path, | ||
| 'success': True, | ||
| 'message': None | ||
| }) | ||
|
|
||
| elif response.status_code == 400: | ||
| _LOGGER.warning( | ||
| "%s teaching of file %s failed with message:%s", | ||
| CLASSIFIER, file_path, response.text) | ||
| self.hass.bus.fire( | ||
| EVENT_CLASSIFIER_TEACH, { | ||
| ATTR_CLASSIFIER: CLASSIFIER, | ||
| ATTR_NAME: name, | ||
| FILE_PATH: file_path, | ||
| 'success': False, | ||
| 'message': response.text | ||
| }) | ||
| teach_file( | ||
| self._url_teach, self._username, self._password, name, file_path) | ||
|
|
||
| @property | ||
| def camera_entity(self): | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -7,14 +7,13 @@ | |
|
|
||
| from homeassistant.core import callback | ||
| from homeassistant.const import ( | ||
| ATTR_ENTITY_ID, ATTR_NAME, CONF_FRIENDLY_NAME, | ||
| CONF_IP_ADDRESS, CONF_PORT, STATE_UNKNOWN) | ||
| ATTR_ENTITY_ID, ATTR_NAME, CONF_FRIENDLY_NAME, CONF_PASSWORD, | ||
| CONF_USERNAME, CONF_IP_ADDRESS, CONF_PORT, | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. trailing whitespace |
||
| HTTP_BAD_REQUEST, HTTP_UNAUTHORIZED, HTTP_OK, STATE_UNKNOWN) | ||
| from homeassistant.setup import async_setup_component | ||
| import homeassistant.components.image_processing as ip | ||
| import homeassistant.components.image_processing.facebox as fb | ||
|
|
||
| # pylint: disable=redefined-outer-name | ||
|
|
||
| MOCK_IP = '192.168.0.1' | ||
| MOCK_PORT = '8080' | ||
|
|
||
|
|
@@ -33,9 +32,11 @@ | |
| "faces": [MOCK_FACE]} | ||
|
|
||
| MOCK_NAME = 'mock_name' | ||
| MOCK_USERNAME = 'mock_username' | ||
| MOCK_PASSWORD = 'mock_password' | ||
|
|
||
| # Faces data after parsing. | ||
| PARSED_FACES = [{ATTR_NAME: 'John Lennon', | ||
| PARSED_FACES = [{fb.FACEBOX_NAME: 'John Lennon', | ||
| fb.ATTR_IMAGE_ID: 'john.jpg', | ||
| fb.ATTR_CONFIDENCE: 58.12, | ||
| fb.ATTR_MATCHED: True, | ||
|
|
@@ -114,6 +115,17 @@ async def test_setup_platform(hass): | |
| assert hass.states.get(VALID_ENTITY_ID) | ||
|
|
||
|
|
||
| async def test_setup_platform_with_auth(hass): | ||
| """Setup platform with one entity and auth.""" | ||
|
|
||
| valid_config_auth = VALID_CONFIG.copy() | ||
| valid_config_auth[ip.DOMAIN][CONF_USERNAME] = MOCK_USERNAME | ||
| valid_config_auth[ip.DOMAIN][CONF_PASSWORD] = MOCK_PASSWORD | ||
|
|
||
| await async_setup_component(hass, ip.DOMAIN, valid_config_auth) | ||
| assert hass.states.get(VALID_ENTITY_ID) | ||
|
|
||
|
|
||
| async def test_process_image(hass, mock_image): | ||
| """Test processing of an image.""" | ||
| await async_setup_component(hass, ip.DOMAIN, VALID_CONFIG) | ||
|
|
@@ -183,40 +195,23 @@ async def test_teach_service(hass, mock_image, mock_isfile, mock_open_file): | |
| await async_setup_component(hass, ip.DOMAIN, VALID_CONFIG) | ||
| assert hass.states.get(VALID_ENTITY_ID) | ||
|
|
||
| teach_events = [] | ||
|
|
||
| @callback | ||
| def mock_teach_event(event): | ||
| """Mock event.""" | ||
| teach_events.append(event) | ||
|
|
||
| hass.bus.async_listen( | ||
| 'image_processing.teach_classifier', mock_teach_event) | ||
|
|
||
| # Patch out 'is_allowed_path' as the mock files aren't allowed | ||
| hass.config.is_allowed_path = Mock(return_value=True) | ||
|
|
||
| with requests_mock.Mocker() as mock_req: | ||
| url = "http://{}:{}/facebox/teach".format(MOCK_IP, MOCK_PORT) | ||
| mock_req.post(url, status_code=200) | ||
| mock_req.post(url, status_code=HTTP_OK) | ||
| data = {ATTR_ENTITY_ID: VALID_ENTITY_ID, | ||
| ATTR_NAME: MOCK_NAME, | ||
| fb.FILE_PATH: MOCK_FILE_PATH} | ||
| await hass.services.async_call( | ||
| ip.DOMAIN, fb.SERVICE_TEACH_FACE, service_data=data) | ||
| await hass.async_block_till_done() | ||
|
|
||
| assert len(teach_events) == 1 | ||
| assert teach_events[0].data[fb.ATTR_CLASSIFIER] == fb.CLASSIFIER | ||
| assert teach_events[0].data[ATTR_NAME] == MOCK_NAME | ||
| assert teach_events[0].data[fb.FILE_PATH] == MOCK_FILE_PATH | ||
| assert teach_events[0].data['success'] | ||
| assert not teach_events[0].data['message'] | ||
|
|
||
| # Now test the failed teaching. | ||
| with requests_mock.Mocker() as mock_req: | ||
| url = "http://{}:{}/facebox/teach".format(MOCK_IP, MOCK_PORT) | ||
| mock_req.post(url, status_code=400, text=MOCK_ERROR) | ||
| mock_req.post(url, status_code=HTTP_BAD_REQUEST, text=MOCK_ERROR) | ||
| data = {ATTR_ENTITY_ID: VALID_ENTITY_ID, | ||
| ATTR_NAME: MOCK_NAME, | ||
| fb.FILE_PATH: MOCK_FILE_PATH} | ||
|
|
@@ -225,13 +220,6 @@ def mock_teach_event(event): | |
| service_data=data) | ||
| await hass.async_block_till_done() | ||
|
|
||
| assert len(teach_events) == 2 | ||
| assert teach_events[1].data[fb.ATTR_CLASSIFIER] == fb.CLASSIFIER | ||
| assert teach_events[1].data[ATTR_NAME] == MOCK_NAME | ||
| assert teach_events[1].data[fb.FILE_PATH] == MOCK_FILE_PATH | ||
| assert not teach_events[1].data['success'] | ||
| assert teach_events[1].data['message'] == MOCK_ERROR | ||
|
|
||
|
|
||
| async def test_setup_platform_with_name(hass): | ||
| """Setup platform with one entity and a name.""" | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
I see this as a platform variable